diff --git a/docs/CN/source/tutorial/api_server_args.rst b/docs/CN/source/tutorial/api_server_args.rst
index 8e7f9d78e8..3417e7c01f 100644
--- a/docs/CN/source/tutorial/api_server_args.rst
+++ b/docs/CN/source/tutorial/api_server_args.rst
@@ -294,6 +294,10 @@ PD 分离模式参数
当输入图片超过该阈值时,LightLLM 会先自动将其缩放到该像素预算内,再继续后续流程。
+.. option:: --disable_image_resize
+
+ 禁用对超过 ``--max_image_pixels`` 的图片的自动缩放。默认开启自动缩放。
+
.. option:: --visual_infer_batch_size
每次推理批次中处理的图像数量,默认为 ``1``
@@ -310,10 +314,6 @@ PD 分离模式参数
ViT 的数据并行实例数量,默认为 ``1``
-.. option:: --visual_nccl_ports
-
- 为 ViT 构建分布式环境的 NCCL 端口列表,例如 29500 29501 29502,默认为 [29500]
-
.. option:: --vit_att_backend
设置 ViT 使用的注意力后端。可选值为:
@@ -496,9 +496,9 @@ PD 分离模式参数
* ``triton``: 使用 torch 和 triton kernel(默认)
* ``sglang_kernel``: 使用 sglang_kernel 实现
-.. option:: --return_all_prompt_logprobs
+.. option:: --enable_prompt_logprobs
- 返回所有提示 token 的 logprobs
+ 启用 prompt top-k logprobs 捕获
.. option:: --use_reward_model
@@ -549,14 +549,6 @@ DeepSeek 冗余专家参数
监控和日志参数
--------------
-.. option:: --disable_log_stats
-
- 禁用吞吐量统计日志记录
-
-.. option:: --log_stats_interval
-
- 记录统计信息的间隔(秒),默认为 ``10``
-
.. option:: --health_monitor
检查服务健康状态并在出错时重启
diff --git a/docs/EN/source/tutorial/api_server_args.rst b/docs/EN/source/tutorial/api_server_args.rst
index 84785de3b7..6cd9252b29 100644
--- a/docs/EN/source/tutorial/api_server_args.rst
+++ b/docs/EN/source/tutorial/api_server_args.rst
@@ -293,6 +293,10 @@ Multimodal Parameters
If an input image exceeds this threshold, LightLLM automatically resizes it down to this pixel budget before continuing.
+.. option:: --disable_image_resize
+
+ Disable automatic resize for images exceeding ``--max_image_pixels``. Resize is enabled by default.
+
.. option:: --visual_infer_batch_size
Number of images processed in each inference batch, default is ``1``
@@ -309,10 +313,6 @@ Multimodal Parameters
Number of data parallel instances for ViT, default is ``1``
-.. option:: --visual_nccl_ports
-
- List of NCCL ports for ViT, e.g., 29500 29501 29502, default is [29500]
-
.. option:: --vit_att_backend
Set the attention backend for ViT. Available options:
@@ -497,9 +497,9 @@ Sampling and Generation Parameters
* ``triton``: Use torch and triton kernel (default)
* ``sglang_kernel``: Use sglang_kernel implementation
-.. option:: --return_all_prompt_logprobs
+.. option:: --enable_prompt_logprobs
- Return logprobs for all prompt tokens
+ Enable prompt top-k logprobs capture
.. option:: --use_reward_model
@@ -550,14 +550,6 @@ DeepSeek Redundant Expert Parameters
Monitoring and Logging Parameters
---------------------------------
-.. option:: --disable_log_stats
-
- Disable throughput statistics logging
-
-.. option:: --log_stats_interval
-
- Interval for recording statistics (seconds), default is ``10``
-
.. option:: --health_monitor
Check service health status and restart on error
diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py
index 779ce454cc..72d3f26cd9 100755
--- a/lightllm/common/basemodel/basemodel.py
+++ b/lightllm/common/basemodel/basemodel.py
@@ -36,6 +36,10 @@
)
from lightllm.common.triton_utils.autotuner import Autotuner
from lightllm.utils.infer_utils import post_empty_cache
+from lightllm.utils.torch_memory_saver_utils import (
+ TorchMemorySaverWrapper,
+ MemoryTag,
+)
from .attention import get_prefill_att_backend_class, get_decode_att_backend_class
from .attention import BaseAttBackend
@@ -93,6 +97,7 @@ def __init__(self, kvargs):
self.tp_world_size_ = get_dp_world_size()
self.enable_tpsp_mix_mode = get_env_start_args().enable_tpsp_mix_mode
+ self.torch_memory_saver = TorchMemorySaverWrapper(self.args.enable_torch_memory_saver)
self.is_mtp_mode = self.args.mtp_mode in [
"vanilla_with_att",
"eagle_with_att",
@@ -106,18 +111,21 @@ def __init__(self, kvargs):
self._verify_params()
self._init_quant()
- self._init_weights()
- self._init_req_manager()
- self._init_mem_manager()
+ enable_weight_cpu_backup = self.args.enable_weight_cpu_backup
+ with self.torch_memory_saver.region(tag=MemoryTag.WEIGHT, enable_cpu_backup=enable_weight_cpu_backup):
+ self._init_weights()
+ with self.torch_memory_saver.region(tag=MemoryTag.KV_CACHE):
+ self._init_req_manager()
+ self._init_mem_manager()
+
# 因为类似 qwen3.5 的linear 架构的模型,其 req_manager 会存储运行时使用的大量 linear state
# 这可能会占用大量的显存,所以,req_manger 中保存的 mem_manger 是mem manager 初始化后再赋值
self.req_manager.mem_manager = self.mem_manager
-
self._check_mem_size()
self._init_infer_layer()
self._init_some_value()
self._init_custom()
- self._load_hf_weights()
+ self.load_weights(self.weight_dict)
self._init_att_backend()
self._init_att_backend1()
@@ -176,13 +184,14 @@ def _init_weights(self, start_layer_index=0):
]
return
- def _load_hf_weights(self):
+ def load_weights(self, weight_dict: dict):
+ assert weight_dict is None or isinstance(weight_dict, dict), "weight_dict must be a dict or None"
load_hf_weights(
- self.data_type,
+ data_type=self.data_type,
weight_dir=self.weight_dir_,
pre_post_layer=self.pre_post_weight,
transformer_layer_list=self.trans_layers_weight,
- weight_dict=self.weight_dict,
+ weight_dict=weight_dict,
)
self.pre_post_weight.verify_load()
[weight.verify_load() for weight in self.trans_layers_weight]
@@ -520,13 +529,13 @@ def _create_unpad_decode_model_output(self, model_output: ModelOutput, origin_ba
def _create_unpad_prefill_model_output(
self, padded_model_output: ModelOutput, origin_handle_token_num: int, origin_batch_size: int
):
- if self.return_all_prompt_logics:
- new_model_output = copy.copy(padded_model_output)
- new_model_output.logits = new_model_output.logits[0:origin_handle_token_num]
- else:
- new_model_output = copy.copy(padded_model_output)
- # 移除多余的pad 的那个 req 对应的 logics
- new_model_output.logits = new_model_output.logits[0:origin_batch_size]
+ new_model_output = copy.copy(padded_model_output)
+ # logits 始终只对应每个请求最后一个位置,移除 padding 的 req 对应的行。
+ new_model_output.logits = new_model_output.logits[0:origin_batch_size]
+ # prompt_logics 保存整个 prefill 阶段所有 token 位置的 logits,
+ # 按实际处理的 token 数量裁剪掉 padding 部分(仅 return_all_prompt_logics 模式下非空)。
+ if new_model_output.prompt_logics is not None:
+ new_model_output.prompt_logics = new_model_output.prompt_logics[0:origin_handle_token_num]
# 特殊模型,特殊模式的特殊变量的特殊 unpad
if new_model_output.mtp_main_output_hiddens is not None:
@@ -701,7 +710,7 @@ def prefill_func(input_tensors, infer_state):
last_input_embs = infer_state._all_to_all_unbalance_get(data=last_input_embs)
predict_logits = self.post_infer.token_forward(last_input_embs, infer_state, self.pre_post_weight)
- model_output = ModelOutput(logits=predict_logits)
+ model_output = ModelOutput(logits=predict_logits, prompt_logics=infer_state.prompt_logics)
# 特殊模型特殊模式的额外输出
if self.is_mtp_mode:
@@ -974,8 +983,8 @@ def _overlap_tpsp_context_forward(self, infer_state: InferStateInfo, infer_state
)
g_cache_manager.cache_env_out()
- model_output = ModelOutput(logits=predict_logits.contiguous())
- model_output1 = ModelOutput(logits=predict_logits1.contiguous())
+ model_output = ModelOutput(logits=predict_logits.contiguous(), prompt_logics=infer_state.prompt_logics)
+ model_output1 = ModelOutput(logits=predict_logits1.contiguous(), prompt_logics=infer_state1.prompt_logics)
if self.is_mtp_mode:
input_embs = self.pre_infer._tpsp_allgather(input=input_embs, infer_state=infer_state)
@@ -1087,6 +1096,7 @@ def _check_max_len_infer(self):
)
logger.error(exception_str)
raise Exception(exception_str)
+ torch.cuda.empty_cache()
return
def autotune_layers(self):
@@ -1221,6 +1231,9 @@ def _init_padded_req(self):
del b_seq_len
del b_ready_cache_len
del model_output
+ del b_mtp_index
+ del b_prefill_start_loc
+ del b_q_seq_len
torch.cuda.empty_cache()
return
diff --git a/lightllm/common/basemodel/batch_objs.py b/lightllm/common/basemodel/batch_objs.py
index 303e6d06ce..28b617a0ec 100644
--- a/lightllm/common/basemodel/batch_objs.py
+++ b/lightllm/common/basemodel/batch_objs.py
@@ -112,6 +112,12 @@ class ModelOutput:
# 输入
mtp_main_output_hiddens: Optional[torch.Tensor] = None
+ # prompt_logics 用于在开启 return_all_prompt_logics 模式(如 enable_prompt_logprobs)时,
+ # 保存整个 prefill 阶段每一个 token 位置对应的 logits(而非仅最后一个位置的 logits)。
+ # 此时 logits 依然只保存每个请求最后一个位置的 logits,prompt_logics 为可选项,仅在
+ # 需要返回 prompt logprobs 信息时才会非空。
+ prompt_logics: Optional[torch.Tensor] = None
+
def to_no_ref_tensor(self):
self.logits = tensor_to_no_ref_tensor(self.logits)
if self.mtp_main_output_hiddens is not None:
diff --git a/lightllm/common/basemodel/cuda_graph.py b/lightllm/common/basemodel/cuda_graph.py
index 6268d6fc35..6c8bf37c78 100644
--- a/lightllm/common/basemodel/cuda_graph.py
+++ b/lightllm/common/basemodel/cuda_graph.py
@@ -8,6 +8,10 @@
from lightllm.utils.envs_utils import get_env_start_args
from lightllm.distributed import dist_group_manager
from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput
+from lightllm.utils.torch_memory_saver_utils import (
+ TorchMemorySaverWrapper,
+ MemoryTag,
+)
from .infer_struct import InferStateInfo
@@ -54,6 +58,7 @@ def __init__(self, max_batch_size=8, max_len_in_batch=8192, tp_world_size: int =
self.max_batch_size = max_batch_size
self.graph_max_len_in_batch = max_len_in_batch
self.enable_decode_microbatch_overlap = self.args.enable_decode_microbatch_overlap
+ self.torch_memory_saver = TorchMemorySaverWrapper(self.args.enable_torch_memory_saver)
self.cuda_graph_batch_sizes = self.gen_cuda_graph_batch_sizes(
max_batch_size=max_batch_size,
@@ -105,7 +110,7 @@ def _capture_decode(self, decode_func, infer_state: InferStateInfo):
if param_name not in pure_para_set:
delattr(infer_state, param_name)
- with torch.cuda.graph(graph_obj, pool=self.mempool):
+ with self.torch_memory_saver.cuda_graph(graph_obj, pool=self.mempool):
model_output = decode_func(infer_state)
self.graph[batch_size] = (graph_obj, infer_state, model_output)
graph_obj.replay()
@@ -139,7 +144,7 @@ def _capture_decode_overlap(
if para_name not in pure_para_set1:
delattr(infer_state1, para_name)
- with torch.cuda.graph(graph_obj, pool=self.mempool):
+ with self.torch_memory_saver.cuda_graph(graph_obj, pool=self.mempool):
model_output, model_output1 = decode_func(infer_state, infer_state1)
self.graph[batch_size] = (
graph_obj,
diff --git a/lightllm/common/basemodel/infer_struct.py b/lightllm/common/basemodel/infer_struct.py
index 89e0508608..10c35759aa 100755
--- a/lightllm/common/basemodel/infer_struct.py
+++ b/lightllm/common/basemodel/infer_struct.py
@@ -56,6 +56,10 @@ def __init__(self):
self.is_token_healing: bool = False
self.return_all_prompt_logics: bool = False
+ # 在开启 return_all_prompt_logics 模式时,保存整个 prefill 阶段每一个
+ # token 位置的 logits,供后续回传 prompt logprobs 信息使用。
+ # 仅在 prefill 阶段且需要返回 prompt logprobs 时才会被填充。
+ self.prompt_logics: Optional[torch.Tensor] = None
self.multimodal_params: dict = None
self.is_cuda_graph: bool = False # 标记是否是cuda graph的捕获推理
self.dist_group: CustomProcessGroup = None
diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/__init__.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/__init__.py
new file mode 100644
index 0000000000..e69de29bb2
diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py
index 26f2b338b7..c20acb12f7 100644
--- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py
+++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py
@@ -9,6 +9,7 @@
SliceMixinTpl,
)
from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.impl import select_fuse_moe_impl
+from lightllm.common.basemodel.moe_route_info_manager import get_moe_capture_callback
from lightllm.common.quantization.quantize_method import QuantizationMethod
from lightllm.utils.envs_utils import get_redundancy_expert_ids, get_redundancy_expert_num, get_env_start_args
from lightllm.utils.dist_utils import get_global_world_size, get_global_rank
@@ -134,9 +135,11 @@ def experts(
topk_group: int,
num_expert_group: int,
is_prefill: Optional[bool] = None,
+ infer_state=None,
shared_expert_gate: Optional[torch.Tensor] = None,
) -> torch.Tensor:
- """Backward compatible method that routes to platform-specific implementation."""
+ # Captures MoE topk expert ids for routed-experts metadata when enabled.
+ moe_capture_callback = get_moe_capture_callback(infer_state, self.layer_num_)
return self.fuse_moe_impl(
input_tensor=input_tensor,
router_logits=router_logits,
@@ -150,6 +153,7 @@ def experts(
topk_group=topk_group,
num_expert_group=num_expert_group,
is_prefill=is_prefill,
+ moe_capture_callback=moe_capture_callback,
per_expert_scale=self.per_expert_scale,
shared_expert_gate=shared_expert_gate,
)
@@ -319,6 +323,7 @@ def _create_weight(self):
device_id=self.device_id_,
num_experts=self.local_n_routed_experts,
)
+ self.w1, self.w3 = w13_param_list
self.w1_list: List[WeightPack] = self._get_expert_weight_list(w13_param_list[0])
self.w3_list: List[WeightPack] = self._get_expert_weight_list(w13_param_list[1])
self.w2_list: List[WeightPack] = self._get_expert_weight_list(self.w2)
@@ -341,7 +346,6 @@ def _get_expert_weight_list(self, weight_pack: WeightPack):
return weight_list
def _load_weight(self, expert_idx_to_local_idx: Dict[int, int], weights: Dict[str, torch.Tensor]):
- # Load each expert with TP slicing
for expert_idx, local_expert_idx in expert_idx_to_local_idx.items():
with self.lock:
self._load_expert(expert_idx, local_expert_idx, weights)
diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/gpt_oss_fused_moe_weight_tp.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/gpt_oss_fused_moe_weight_tp.py
index 90ce5761c3..240bc726ca 100644
--- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/gpt_oss_fused_moe_weight_tp.py
+++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/gpt_oss_fused_moe_weight_tp.py
@@ -4,6 +4,7 @@
from typing import Optional, Tuple, List, Dict, Any
from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.fused_moe_weight import FusedMoeWeight
+from lightllm.common.basemodel.moe_route_info_manager import get_moe_capture_callback
from lightllm.utils.dist_utils import get_current_rank_in_dp, get_current_device_id
from lightllm.common.quantization import Quantcfg
from lightllm.common.quantization.quantize_method import QuantizationMethod
@@ -144,12 +145,18 @@ def experts(
topk_group: int,
num_expert_group: int,
is_prefill: Optional[bool] = None,
+ infer_state=None,
shared_expert_gate: Optional[torch.Tensor] = None,
):
assert shared_expert_gate is None, "shared_expert_gate is not supported by GPT-OSS fused MoE"
topk_weights, topk_ids = self._router(router_logits, top_k)
+ # Captures MoE topk expert ids for routed-experts metadata when enabled.
+ moe_capture_callback = get_moe_capture_callback(infer_state, self.layer_num_)
+ if moe_capture_callback is not None:
+ moe_capture_callback(topk_ids)
+
w1, w1_scale = self.w1
w2, w2_scale = self.w2
use_fp8_w8a8 = self.quant_method is not None
diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/base_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/base_impl.py
index 8467c328da..1e3ad4b196 100644
--- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/base_impl.py
+++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/base_impl.py
@@ -1,10 +1,10 @@
import torch
from abc import abstractmethod
+from typing import Callable, Optional
from lightllm.common.quantization.quantize_method import (
WeightPack,
QuantizationMethod,
)
-from typing import Optional
from lightllm.utils.dist_utils import (
get_global_rank,
get_global_world_size,
@@ -62,6 +62,8 @@ def __call__(
topk_group: int,
num_expert_group: int,
is_prefill: Optional[bool] = None,
+ # Callback to capture MoE topk expert ids (routed experts metadata).
+ moe_capture_callback: Optional[Callable[[torch.Tensor], None]] = None,
per_expert_scale: Optional[torch.Tensor] = None,
# Qwen3.5 uses this gate to control fused shared expert aggregation weights.
shared_expert_gate: Optional[torch.Tensor] = None,
diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py
index 5da17c57e1..aa080a19fe 100644
--- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py
+++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py
@@ -58,7 +58,10 @@ def _select_experts(
topk_weights.mul_(self.routed_scaling_factor)
if per_expert_scale is not None:
topk_weights = topk_weights * per_expert_scale[topk_ids.to(torch.long)].to(topk_weights.dtype)
+ origin_topk_ids = topk_ids
if self.redundancy_expert_num > 0:
+ # 因为 redundancy_topk_ids_repair 会修改 topk_ids,所以需要先复制一份
+ origin_topk_ids = topk_ids.clone()
redundancy_topk_ids_repair(
topk_ids=topk_ids,
redundancy_expert_ids=self.redundancy_expert_ids_tensor,
@@ -67,7 +70,7 @@ def _select_experts(
expert_counter=self.routed_expert_counter_tensor,
enable_counter=self.auto_update_redundancy_expert,
)
- return topk_weights, topk_ids
+ return topk_weights, topk_ids, origin_topk_ids
def _fused_experts(
self,
@@ -104,7 +107,7 @@ def low_latency_dispatch(
n_group: int,
scoring_func: str,
):
- topk_weights, topk_idx = self._select_experts(
+ topk_weights, topk_idx, _ = self._select_experts(
input_tensor=hidden_states,
router_logits=router_logits,
correction_bias=e_score_correction_bias,
diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py
index 110a83094b..1d6a38c069 100644
--- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py
+++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py
@@ -1,5 +1,5 @@
import torch
-from typing import Optional
+from typing import Callable, Optional
from lightllm.common.quantization.no_quant import WeightPack
from lightllm.common.quantization.quantize_method import QuantizationMethod
from .base_impl import FuseMoeBaseImpl
@@ -63,6 +63,7 @@ def _select_experts(
topk_weights.mul_(self.routed_scaling_factor)
if per_expert_scale is not None:
topk_weights = topk_weights * per_expert_scale[topk_ids.to(torch.long)].to(topk_weights.dtype)
+ origin_topk_ids = topk_ids
if self.num_fused_shared_experts > 0:
from lightllm.common.basemodel.triton_kernel.fused_moe.append_shared_expert_topk import (
append_fused_shared_experts,
@@ -75,7 +76,7 @@ def _select_experts(
num_fused_shared_experts=self.num_fused_shared_experts,
shared_expert_gate=shared_expert_gate,
)
- return topk_weights, topk_ids
+ return topk_weights, topk_ids, origin_topk_ids
def _fused_experts(
self,
@@ -120,10 +121,12 @@ def __call__(
topk_group: int,
num_expert_group: int,
is_prefill: Optional[bool] = None,
+ # Callback to capture MoE topk expert ids (routed experts metadata).
+ moe_capture_callback: Optional[Callable[[torch.Tensor], None]] = None,
per_expert_scale: Optional[torch.Tensor] = None,
shared_expert_gate: Optional[torch.Tensor] = None,
):
- topk_weights, topk_ids = self._select_experts(
+ topk_weights, topk_ids, origin_topk_ids = self._select_experts(
input_tensor=input_tensor,
router_logits=router_logits,
correction_bias=correction_bias,
@@ -136,6 +139,10 @@ def __call__(
per_expert_scale=per_expert_scale,
shared_expert_gate=shared_expert_gate,
)
+
+ if moe_capture_callback is not None:
+ moe_capture_callback(origin_topk_ids)
+
output = self._fused_experts(
input_tensor=input_tensor,
w13=w13,
diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/mm_weight/mm_slicer.py b/lightllm/common/basemodel/layer_weights/meta_weights/mm_weight/mm_slicer.py
index ddbf98a866..067c1c8ca9 100644
--- a/lightllm/common/basemodel/layer_weights/meta_weights/mm_weight/mm_slicer.py
+++ b/lightllm/common/basemodel/layer_weights/meta_weights/mm_weight/mm_slicer.py
@@ -28,6 +28,10 @@ def _get_slice_start_end(self, size: int) -> Tuple[int, int]:
end = start + tp_size
return start, end
+ def _assert_weight_ndim(self, tensor: torch.Tensor) -> None:
+ # 2D: 普通 linear (out, in); 3D: MoE 合并权重 (num_experts, out, in)。
+ assert tensor.dim() in (2, 3), f"expect weight ndim in (2, 3), got shape {tuple(tensor.shape)}"
+
class SliceMixinTpl(SliceMixinBase):
def __init__(self, tp_rank: int = None, tp_world_size: int = None, repeat_times: int = 1):
@@ -46,18 +50,20 @@ def _slice_weight_zero_point(self, weight_zero_point: torch.Tensor) -> torch.Ten
raise NotImplementedError("slice_weight_zero_point must implement this method")
-# 默认weight 的shape是 outxin,这也是目前最通用的约定。
-# 所以row-wise是沿着dim=0进行切分,col-wise是沿着dim=1进行切分。
+# 默认 weight 的 shape 末两维是 (out, in),普通 linear 是 2D (out, in),
+# MoE 合并权重则是 3D (num_experts, out, in),统一通过 `...` 处理任意前导维。
+# 约定 row-wise 沿着 out 维(倒数第二维)切分,col-wise 沿着 in 维(最后一维)切分。
class RowSliceMixin(SliceMixinTpl):
def __init__(self, tp_rank: int = None, tp_world_size: int = None, repeat_times: int = 1):
super().__init__(tp_rank, tp_world_size, repeat_times)
def _slice_weight(self, weight: torch.Tensor) -> torch.Tensor:
+ self._assert_weight_ndim(weight)
assert (
- weight.shape[0] * self.repeat_times_ % self.tp_world_size_ == 0
- ), f"tp slice error {weight.shape[0] * self.repeat_times_} % {self.tp_world_size_}"
- start, end = self._get_slice_start_end(weight.shape[0])
- return weight[start:end, :]
+ weight.shape[-2] * self.repeat_times_ % self.tp_world_size_ == 0
+ ), f"tp slice error {weight.shape[-2] * self.repeat_times_} % {self.tp_world_size_}"
+ start, end = self._get_slice_start_end(weight.shape[-2])
+ return weight[..., start:end, :]
def _slice_bias(self, bias: torch.Tensor) -> torch.Tensor:
assert (
@@ -74,18 +80,20 @@ def __init__(self, tp_rank: int = None, tp_world_size: int = None, repeat_times:
super().__init__(tp_rank, tp_world_size, repeat_times)
def _slice_weight_scale(self, weight_scale: torch.Tensor) -> torch.Tensor:
+ self._assert_weight_ndim(weight_scale)
assert (
- weight_scale.shape[0] % self.tp_world_size_ == 0
- ), f"tp slice error {weight_scale.shape[0]} % {self.tp_world_size_}"
- start, end = self._get_slice_start_end(weight_scale.shape[0])
- return weight_scale[start:end]
+ weight_scale.shape[-2] % self.tp_world_size_ == 0
+ ), f"tp slice error {weight_scale.shape[-2]} % {self.tp_world_size_}"
+ start, end = self._get_slice_start_end(weight_scale.shape[-2])
+ return weight_scale[..., start:end, :]
def _slice_weight_zero_point(self, weight_zero_point: torch.Tensor) -> torch.Tensor:
+ self._assert_weight_ndim(weight_zero_point)
assert (
- weight_zero_point.shape[0] % self.tp_world_size_ == 0
- ), f"tp slice error {weight_zero_point.shape[0]} % {self.tp_world_size_}"
- start, end = self._get_slice_start_end(weight_zero_point.shape[0])
- return weight_zero_point[start:end]
+ weight_zero_point.shape[-2] % self.tp_world_size_ == 0
+ ), f"tp slice error {weight_zero_point.shape[-2]} % {self.tp_world_size_}"
+ start, end = self._get_slice_start_end(weight_zero_point.shape[-2])
+ return weight_zero_point[..., start:end, :]
class ColSliceMixin(SliceMixinTpl):
@@ -93,11 +101,12 @@ def __init__(self, tp_rank: int = None, tp_world_size: int = None, repeat_times:
super().__init__(tp_rank, tp_world_size, repeat_times)
def _slice_weight(self, weight: torch.Tensor) -> torch.Tensor:
+ self._assert_weight_ndim(weight)
assert (
- weight.shape[1] * self.repeat_times_ % self.tp_world_size_ == 0
- ), f"tp slice error {weight.shape[1] * self.repeat_times_ } % {self.tp_world_size_}"
- start, end = self._get_slice_start_end(weight.shape[1])
- return weight[:, start:end]
+ weight.shape[-1] * self.repeat_times_ % self.tp_world_size_ == 0
+ ), f"tp slice error {weight.shape[-1] * self.repeat_times_ } % {self.tp_world_size_}"
+ start, end = self._get_slice_start_end(weight.shape[-1])
+ return weight[..., start:end]
def _slice_bias(self, bias: torch.Tensor) -> torch.Tensor:
return bias / self.tp_world_size_ * self.repeat_times_
@@ -108,18 +117,20 @@ def __init__(self, tp_rank: int = None, tp_world_size: int = None, repeat_times:
super().__init__(tp_rank, tp_world_size, repeat_times)
def _slice_weight_scale(self, weight_scale: torch.Tensor) -> torch.Tensor:
+ self._assert_weight_ndim(weight_scale)
assert (
- weight_scale.shape[1] * self.repeat_times_ % self.tp_world_size_ == 0
- ), f"tp slice error {weight_scale.shape[1] * self.repeat_times_ } % {self.tp_world_size_}"
- start, end = self._get_slice_start_end(weight_scale.shape[1])
- return weight_scale[:, start:end]
+ weight_scale.shape[-1] * self.repeat_times_ % self.tp_world_size_ == 0
+ ), f"tp slice error {weight_scale.shape[-1] * self.repeat_times_ } % {self.tp_world_size_}"
+ start, end = self._get_slice_start_end(weight_scale.shape[-1])
+ return weight_scale[..., start:end]
def _slice_weight_zero_point(self, weight_zero_point: torch.Tensor) -> torch.Tensor:
+ self._assert_weight_ndim(weight_zero_point)
assert (
- weight_zero_point.shape[1] * self.repeat_times_ % self.tp_world_size_ == 0
- ), f"tp slice error {weight_zero_point.shape[1] * self.repeat_times_ } % {self.tp_world_size_}"
- start, end = self._get_slice_start_end(weight_zero_point.shape[1])
- return weight_zero_point[:, start:end]
+ weight_zero_point.shape[-1] * self.repeat_times_ % self.tp_world_size_ == 0
+ ), f"tp slice error {weight_zero_point.shape[-1] * self.repeat_times_ } % {self.tp_world_size_}"
+ start, end = self._get_slice_start_end(weight_zero_point.shape[-1])
+ return weight_zero_point[..., start:end]
# awq 的量化权重是inxout存储格式,需要定制实现。
diff --git a/lightllm/common/basemodel/logprobs_manager.py b/lightllm/common/basemodel/logprobs_manager.py
new file mode 100644
index 0000000000..b905a8dced
--- /dev/null
+++ b/lightllm/common/basemodel/logprobs_manager.py
@@ -0,0 +1,169 @@
+"""Prompt logprobs 捕获与配置管理。
+
+本模块为 ``--enable_prompt_logprobs`` 提供端到端支持:在 prefill 过程中记录
+prompt 各位置的 top-k token id 与对应 logprob,并在请求结束时写入 final token
+metadata shm,供 HTTP 进程编码进 API 返回的 ``prompt_logprobs``。
+
+背景与目标
+----------
+开启 ``enable_prompt_logprobs`` 后,请求可通过 ``prompt_logprobs=k`` 要求返回
+每个 prompt 位置的 top-k 候选。推理侧需要在 prefill 产出 logits 时把 top-k
+结果按 KV 槽位落到 CPU pinned buffer,避免阻塞 GPU,并在请求结束时按 mem
+indexes 导出。
+
+``prompt_logprobs=0`` 走另一条路径(返回真实命中 token 的 logprob/rank),
+不经过本 manager 的 top-k buffer。
+
+核心职责
+--------
+1. **配置(phase-1)**
+ 记录 ``max_topk``(由环境变量 ``LIGHTLLM_MAX_PROMPT_LOGPROBS`` 控制上限)。
+
+2. **捕获缓冲(phase-2,可选)**
+ 在 infer 进程按 KV cache 槽位分配 pinned CPU buffer:
+ ``top_token_ids[kv_slot, max_topk]`` / ``top_logprobs[kv_slot, max_topk]``。
+ Triton kernel 将 GPU 上的 top-k 结果 scatter 写入该 buffer。
+
+3. **捕获 / 导出 / 槽位拷贝**
+ - ``capture``:prefill 时写入对应 mem indexes
+ - ``extract``:请求结束时按 mem indexes 取出 ``(token_ids, logprobs)``
+ - ``copy_slots``:radix cache 等场景下在 KV 槽位间复制已捕获数据
+
+进程与初始化
+------------
+- 进程内单例:``PromptLogprobsCaptureManager.get_instance()``。
+ 仅在 ``enable_prompt_logprobs`` 时创建,且只做 phase-1 配置初始化。
+- **Infer 进程(通常 dp master)**:在 phase-1 之后调用
+ ``init_capture_buffer(kv_cache_size)``,再参与 capture / extract。
+
+数据流简图::
+
+ prefill logits → topk
+ │
+ ▼
+ capture(mem_indexes) ──scatter──► pinned CPU buffers
+ │
+ ▼ (request finished)
+ extract(mem_indexes, topk)
+ │
+ ▼
+ final_token_metadata shm ──HTTP──► response["prompt_logprobs"]
+"""
+
+import os
+from typing import ClassVar, Optional, Tuple
+
+import numpy as np
+import torch
+
+from lightllm.common.basemodel.triton_kernel.logprobs_capture import scatter_prompt_logprobs_to_cpu
+from lightllm.utils.log_utils import init_logger
+
+logger = init_logger(__name__)
+
+_MAX_PROMPT_LOGPROBS = int(os.getenv("LIGHTLLM_MAX_PROMPT_LOGPROBS", 128))
+
+
+class PromptLogprobsCaptureManager:
+ """管理 prompt top-k logprobs 的配置、捕获缓冲与导出。
+
+ 详见模块文档字符串。
+ """
+
+ _instance: ClassVar[Optional["PromptLogprobsCaptureManager"]] = None
+
+ @classmethod
+ def get_instance(cls) -> Optional["PromptLogprobsCaptureManager"]:
+ """Return the process singleton with phase-1 (config) init only.
+
+ Capture buffer is optional and must be allocated separately via
+ ``init_capture_buffer`` when needed (infer process).
+ """
+ if cls._instance is not None:
+ return cls._instance
+
+ from lightllm.utils.envs_utils import get_env_start_args
+
+ args = get_env_start_args()
+ if not args.enable_prompt_logprobs:
+ return None
+
+ cls._instance = cls(max_topk=_MAX_PROMPT_LOGPROBS)
+ return cls._instance
+
+ def __init__(self, max_topk: int):
+ """Phase-1 init: config metadata only. Call init_capture_buffer() when capture is needed."""
+ self.max_topk = max_topk
+ self.kv_cache_size: Optional[int] = None
+ self.top_token_ids: Optional[torch.Tensor] = None
+ self.top_logprobs: Optional[torch.Tensor] = None
+ self.top_token_ids_ptr: Optional[torch.Tensor] = None
+ self.top_logprobs_ptr: Optional[torch.Tensor] = None
+
+ logger.info(f"PromptLogprobsCaptureManager created: max_topk={max_topk}")
+
+ def is_buffer_initialized(self) -> bool:
+ """Whether phase-2 capture buffers have been allocated."""
+ return self.top_token_ids is not None
+
+ def init_capture_buffer(self, kv_cache_size: int) -> None:
+ """Phase-2 init: allocate pinned CPU buffers for capture/extract."""
+ if self.is_buffer_initialized():
+ return
+
+ self.kv_cache_size = kv_cache_size
+ self.top_token_ids = torch.empty(
+ (kv_cache_size, self.max_topk), dtype=torch.int32, device="cpu", pin_memory=True
+ )
+ self.top_logprobs = torch.empty(
+ (kv_cache_size, self.max_topk), dtype=torch.float32, device="cpu", pin_memory=True
+ )
+ self.top_token_ids_ptr = torch.tensor([self.top_token_ids.data_ptr()], dtype=torch.uint64, device="cuda")
+ self.top_logprobs_ptr = torch.tensor([self.top_logprobs.data_ptr()], dtype=torch.uint64, device="cuda")
+
+ pinned_bytes = self.top_token_ids.numel() * self.top_token_ids.element_size()
+ pinned_bytes += self.top_logprobs.numel() * self.top_logprobs.element_size()
+ logger.info(
+ f"PromptLogprobsCaptureManager capture buffer ready: kv_cache_size={kv_cache_size}, "
+ f"max_topk={self.max_topk}, pinned_memory={pinned_bytes / 1024 / 1024:.2f}MB"
+ )
+
+ def capture(
+ self,
+ mem_indexes: torch.Tensor,
+ top_token_ids: torch.Tensor,
+ top_logprobs: torch.Tensor,
+ ) -> None:
+ if not self.is_buffer_initialized():
+ return
+ scatter_prompt_logprobs_to_cpu(
+ mem_indexes=mem_indexes,
+ top_token_ids=top_token_ids,
+ top_logprobs=top_logprobs,
+ top_token_ids_buffer_ptr=self.top_token_ids_ptr,
+ top_logprobs_buffer_ptr=self.top_logprobs_ptr,
+ kv_cache_size=self.kv_cache_size,
+ max_topk=self.max_topk,
+ )
+
+ def extract(self, mem_indexes: torch.Tensor, topk: int) -> Tuple[np.ndarray, np.ndarray]:
+ if not self.is_buffer_initialized():
+ return
+ indexes = mem_indexes.cpu() if mem_indexes.is_cuda else mem_indexes
+ return (
+ self.top_token_ids[indexes, :topk].numpy(),
+ self.top_logprobs[indexes, :topk].numpy(),
+ )
+
+ def copy_slots(
+ self,
+ source_indexes: torch.Tensor,
+ destination_indexes: torch.Tensor,
+ topk: int,
+ ) -> None:
+ if not self.is_buffer_initialized():
+ return
+ source = source_indexes.cpu() if source_indexes.is_cuda else source_indexes
+ destination = destination_indexes.cpu() if destination_indexes.is_cuda else destination_indexes
+ self.top_token_ids[destination, :topk] = self.top_token_ids[source, :topk]
+ self.top_logprobs[destination, :topk] = self.top_logprobs[source, :topk]
diff --git a/lightllm/common/basemodel/moe_route_info_manager.py b/lightllm/common/basemodel/moe_route_info_manager.py
new file mode 100644
index 0000000000..a9732317a3
--- /dev/null
+++ b/lightllm/common/basemodel/moe_route_info_manager.py
@@ -0,0 +1,263 @@
+"""MoE routed-experts 捕获与配置管理。
+
+本模块为 ``--enable_return_routed_experts`` 提供端到端支持:在 MoE 推理过程中记录
+每个 token、每一层选中的 top-k expert id,并在请求结束时随响应返回
+``routed_experts`` 元数据。
+
+背景与目标
+----------
+MoE 模型每个 token 只会激活少量专家。训练、评测、路由分析等场景常需要知道
+「某个 token 在各 MoE 层实际走到了哪些 expert」。LightLLM 在开启
+``enable_return_routed_experts`` 后,由本模块在推理路径上零拷贝地收集这些
+topk ids,最终写入请求的 final token metadata shm,供 HTTP 进程读取并编码进
+API 返回。
+
+非 MoE 模型不应开启该功能;配置解析假定模型具备合法的 MoE 字段
+(专家数、topk、MoE 层分布等)。
+
+核心职责
+--------
+1. **路由配置(phase-1)**
+ 从 ``config.json`` 解析:
+ - ``num_moe_layers`` / ``topk`` / ``dtype_id``
+ - ``layer_index_to_moe_index``:transformer layer index → 稠密 MoE 槽位 index
+ (部分模型前若干层是 dense MLP,或按 ``moe_layer_freq`` /
+ ``decoder_sparse_step`` 稀疏分布 MoE 层)。
+
+2. **捕获缓冲(phase-2,可选)**
+ 在 infer 进程按 KV cache 槽位分配 pinned CPU buffer:
+ ``routing_buffer[kv_slot, moe_layer_slot, topk]``。
+ Triton kernel 将 GPU 上的 topk ids scatter 写入该 buffer,避免同步 D2H。
+
+3. **捕获回调**
+ fused MoE 前向在选出 topk ids 后调用 ``moe_capture_callback``,按当前
+ ``mem_indexes``(token 对应的 KV 槽位)写入对应层的路由信息。
+
+4. **导出**
+ 请求结束时按该请求占用的 mem indexes ``extract`` 出
+ ``(num_tokens, num_moe_layers, topk)`` 数组,写入 final token metadata,
+ HTTP 侧再编码为响应中的 ``routed_experts``。
+
+进程与初始化
+------------
+- 进程内单例:``MoeRouteInfoManager.get_instance()``。
+ 仅在 ``enable_return_routed_experts`` 时创建,且只做 phase-1 配置初始化。
+- **HTTP 进程**:只需配置元数据(层数 / topk / dtype),用于解析 shm 中的
+ routed experts 布局,不分配 capture buffer。
+- **Infer 进程(通常 dp_rank==0)**:在 phase-1 之后调用
+ ``init_capture_buffer(kv_cache_size)``,再参与 capture / extract。
+
+数据流简图::
+
+ MoE forward (topk_ids)
+ │
+ ▼
+ moe_capture_callback ──scatter──► routing_buffer[mem_index, moe_slot, :]
+ │
+ ▼ (request finished)
+ extract(mem_indexes)
+ │
+ ▼
+ final_token_metadata shm ──HTTP──► response["routed_experts"]
+"""
+
+import json
+import os
+import torch
+import numpy as np
+from typing import ClassVar, Dict, Optional, Tuple
+from lightllm.common.basemodel.triton_kernel.routing_capture import scatter_routing_topk_to_cpu
+from lightllm.utils.log_utils import init_logger
+
+logger = init_logger(__name__)
+
+
+class MoeRouteInfoManager:
+ """管理 MoE topk expert id 的配置、捕获缓冲与导出。
+
+ 详见模块文档字符串。
+ """
+
+ _instance: ClassVar[Optional["MoeRouteInfoManager"]] = None
+
+ @classmethod
+ def get_instance(cls) -> Optional["MoeRouteInfoManager"]:
+ """Return the process singleton with phase-1 (route config) init only.
+
+ Capture buffer is optional and must be allocated separately via
+ ``init_capture_buffer`` when needed (infer process).
+ """
+ if cls._instance is not None:
+ return cls._instance
+
+ from lightllm.utils.envs_utils import get_env_start_args
+
+ args = get_env_start_args()
+ if not args.enable_return_routed_experts:
+ return None
+
+ num_moe_layers, topk, dtype_id, layer_index_to_moe_index = cls.get_route_config_from_model_dir(args.model_dir)
+ cls._instance = cls(
+ num_moe_layers=num_moe_layers,
+ topk=topk,
+ dtype_id=dtype_id,
+ layer_index_to_moe_index=layer_index_to_moe_index,
+ )
+ return cls._instance
+
+ @staticmethod
+ def _get_layer_index_to_moe_index_from_config(config: dict) -> Dict[int, int]:
+ """Build layer_index -> dense moe-slot index from model config."""
+ num_layers = config.get("num_hidden_layers", config.get("n_layer", config.get("num_layers", 0)))
+ num_experts = config.get("n_routed_experts", config.get("num_experts", config.get("num_local_experts", 0)))
+ assert num_layers > 0 and num_experts > 0
+
+ if "first_k_dense_replace" in config:
+ first_k_dense_replace = config.get("first_k_dense_replace", 0)
+ moe_layer_freq = config.get("moe_layer_freq", 1)
+ moe_layer_indexes = [
+ layer_index
+ for layer_index in range(num_layers)
+ if layer_index >= first_k_dense_replace and layer_index % moe_layer_freq == 0
+ ]
+ elif "mlp_only_layers" in config or "decoder_sparse_step" in config:
+ mlp_only_layers = set(config.get("mlp_only_layers", []))
+ decoder_sparse_step = config.get("decoder_sparse_step", 1)
+ moe_layer_indexes = [
+ layer_index
+ for layer_index in range(num_layers)
+ if layer_index not in mlp_only_layers and (layer_index + 1) % decoder_sparse_step == 0
+ ]
+ elif config.get("enable_moe_block", False):
+ moe_layer_indexes = list(range(num_layers))
+ else:
+ moe_layer_indexes = list(range(num_layers))
+
+ assert len(moe_layer_indexes) > 0
+ return {layer_index: moe_index for moe_index, layer_index in enumerate(moe_layer_indexes)}
+
+ @staticmethod
+ def get_route_config_from_model_dir(model_dir: str) -> Tuple[int, int, int, Dict[int, int]]:
+ """Return (num_moe_layers, topk, dtype_id, layer_index_to_moe_index) from model config.
+
+ Caller must only use this when --enable_return_routed_experts is set on a MoE model.
+ """
+ with open(os.path.join(model_dir, "config.json"), "r") as json_file:
+ config = json.load(json_file)
+ config = config.get("text_config", config)
+
+ layer_index_to_moe_index = MoeRouteInfoManager._get_layer_index_to_moe_index_from_config(config)
+ num_moe_layers = len(layer_index_to_moe_index)
+ topk = config.get("num_experts_per_tok", config.get("top_k_experts", 0))
+ num_experts = config.get("n_routed_experts", config.get("num_experts", config.get("num_local_experts", 0)))
+ assert topk > 0 and num_experts > 0
+
+ dtype_id = 1 if num_experts <= 256 else 2
+ return num_moe_layers, topk, dtype_id, layer_index_to_moe_index
+
+ def __init__(
+ self,
+ num_moe_layers: int,
+ topk: int,
+ dtype_id: int,
+ layer_index_to_moe_index: Optional[Dict[int, int]] = None,
+ ):
+ """Phase-1 init: route config metadata only. Call init_capture_buffer() when capture is needed."""
+ self.num_moe_layers = num_moe_layers
+ self.topk = topk
+ self.dtype_id = dtype_id
+ self.layer_index_to_moe_index = layer_index_to_moe_index or {i: i for i in range(num_moe_layers)}
+
+ self.kv_cache_size: Optional[int] = None
+ self.routing_buffer: Optional[torch.Tensor] = None
+ self.routing_buffer_ptr: Optional[torch.Tensor] = None
+
+ logger.info(f"MoeRouteInfoManager created: num_moe_layers={num_moe_layers}, topk={topk}, dtype_id={dtype_id}")
+
+ def get_np_dtype(self):
+ if self.dtype_id == 1:
+ return np.uint8
+ elif self.dtype_id == 2:
+ return np.int16
+ return np.int32
+
+ def get_torch_dtype(self):
+ if self.dtype_id == 1:
+ return torch.uint8
+ elif self.dtype_id == 2:
+ return torch.int16
+ return torch.int32
+
+ def is_buffer_initialized(self) -> bool:
+ """Whether phase-2 capture buffer has been allocated."""
+ return self.routing_buffer is not None
+
+ def init_capture_buffer(self, kv_cache_size: int) -> None:
+ """Phase-2 init: allocate pinned CPU routing buffer for capture/extract."""
+ if self.is_buffer_initialized():
+ return
+
+ torch_dtype = self.get_torch_dtype()
+ dtype_bytes = torch_dtype.itemsize
+ # Shape: (kv_cache_size, num_moe_layers, topk). Pinned CPU memory saves GPU memory
+ # while allowing the Triton scatter kernel to write without a synchronous D2H copy.
+ buffer_bytes = self.num_moe_layers * kv_cache_size * self.topk * dtype_bytes
+ self.kv_cache_size = kv_cache_size
+ self.routing_buffer = torch.zeros(
+ (kv_cache_size, self.num_moe_layers, self.topk),
+ dtype=torch_dtype,
+ device="cpu",
+ pin_memory=True,
+ )
+ self.routing_buffer_ptr = torch.tensor([self.routing_buffer.data_ptr()], dtype=torch.uint64, device="cuda")
+
+ logger.info(
+ f"MoeRouteInfoManager capture buffer ready: kv_cache_size={kv_cache_size}, "
+ f"routing_buffer(cpu)={buffer_bytes / 1024 / 1024:.2f}MB, dtype={torch_dtype}"
+ )
+
+ def get_moe_capture_callback(self, layer_index: int, mem_indexes: torch.Tensor):
+ """Return a callback that captures MoE topk expert ids into the routing buffer."""
+ if not self.is_buffer_initialized():
+ return None
+ moe_layer_index = self.layer_index_to_moe_index.get(layer_index)
+ if moe_layer_index is None:
+ return None
+ if not mem_indexes.is_cuda:
+ mem_indexes = mem_indexes.cuda(non_blocking=True)
+
+ # Captures MoE topk ids for this layer and scatters them into the pinned CPU buffer.
+ def moe_capture_callback(topk_ids: torch.Tensor) -> None:
+ self.capture(moe_layer_index=moe_layer_index, topk_ids=topk_ids, mem_indexes=mem_indexes)
+
+ return moe_capture_callback
+
+ def capture(self, moe_layer_index: int, topk_ids: torch.Tensor, mem_indexes: torch.Tensor) -> None:
+ if not self.is_buffer_initialized():
+ return
+ assert topk_ids.dim() == 2
+ assert topk_ids.shape[1] == self.topk
+ assert mem_indexes.shape[0] >= topk_ids.shape[0]
+ scatter_routing_topk_to_cpu(
+ topk_ids=topk_ids,
+ mem_indexes=mem_indexes,
+ routing_buffer_ptr=self.routing_buffer_ptr,
+ moe_layer_index=moe_layer_index,
+ num_moe_layers=self.num_moe_layers,
+ topk=self.topk,
+ dtype_id=self.dtype_id,
+ )
+
+ def extract(self, mem_indexes: torch.Tensor) -> np.ndarray:
+ if not self.is_buffer_initialized():
+ return
+ cpu_indexes = mem_indexes.cpu() if mem_indexes.is_cuda else mem_indexes
+ return self.routing_buffer[cpu_indexes, :, :].numpy()
+
+
+def get_moe_capture_callback(infer_state, layer_index: int):
+ """Return a callback that captures MoE topk expert ids, or None if capture is disabled."""
+ mgr = MoeRouteInfoManager.get_instance()
+ if mgr is None or not mgr.is_buffer_initialized():
+ return None
+ return mgr.get_moe_capture_callback(layer_index=layer_index, mem_indexes=infer_state.mem_index)
diff --git a/lightllm/common/basemodel/triton_kernel/logprobs_capture.py b/lightllm/common/basemodel/triton_kernel/logprobs_capture.py
new file mode 100644
index 0000000000..701ac848ac
--- /dev/null
+++ b/lightllm/common/basemodel/triton_kernel/logprobs_capture.py
@@ -0,0 +1,69 @@
+import torch
+import triton
+import triton.language as tl
+
+
+@triton.jit
+def _scatter_prompt_logprobs_to_cpu(
+ mem_indexes,
+ top_token_ids,
+ top_logprobs,
+ top_token_ids_buffer_ptr,
+ top_logprobs_buffer_ptr,
+ token_count,
+ kv_cache_size,
+ TOPK: tl.constexpr,
+ MAX_TOPK: tl.constexpr,
+ BLOCK: tl.constexpr,
+):
+ offsets = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK)
+ rows = offsets // TOPK
+ columns = offsets - rows * TOPK
+ mask = rows < token_count
+ mem_index = tl.load(mem_indexes + rows, mask=mask, other=-1).to(tl.int64)
+ write_mask = mask & (mem_index >= 0) & (mem_index < kv_cache_size)
+ source_offset = rows * TOPK + columns
+ destination_offset = mem_index * MAX_TOPK + columns
+
+ token_ids = tl.load(top_token_ids + source_offset, mask=write_mask, other=-1)
+ token_ids_dst = tl.load(top_token_ids_buffer_ptr).to(tl.pointer_type(tl.int32))
+ tl.store(token_ids_dst + destination_offset, token_ids, mask=write_mask)
+
+ logprobs = tl.load(top_logprobs + source_offset, mask=write_mask, other=0.0)
+ logprobs_dst = tl.load(top_logprobs_buffer_ptr).to(tl.pointer_type(tl.float32))
+ tl.store(logprobs_dst + destination_offset, logprobs, mask=write_mask)
+
+
+def scatter_prompt_logprobs_to_cpu(
+ mem_indexes: torch.Tensor,
+ top_token_ids: torch.Tensor,
+ top_logprobs: torch.Tensor,
+ top_token_ids_buffer_ptr: torch.Tensor,
+ top_logprobs_buffer_ptr: torch.Tensor,
+ kv_cache_size: int,
+ max_topk: int,
+) -> None:
+ token_count, topk = top_token_ids.shape
+ assert mem_indexes.is_cuda and mem_indexes.is_contiguous()
+ assert mem_indexes.dtype in (torch.int32, torch.int64)
+ assert mem_indexes.numel() == token_count
+ assert top_token_ids.is_cuda and top_token_ids.is_contiguous()
+ assert top_token_ids.dtype == torch.int32
+ assert top_logprobs.is_cuda and top_logprobs.is_contiguous()
+ assert top_logprobs.dtype == torch.float32
+ assert top_logprobs.shape == top_token_ids.shape
+ assert 0 < topk <= max_topk
+
+ block = 1024
+ _scatter_prompt_logprobs_to_cpu[(triton.cdiv(token_count * topk, block),)](
+ mem_indexes=mem_indexes,
+ top_token_ids=top_token_ids,
+ top_logprobs=top_logprobs,
+ top_token_ids_buffer_ptr=top_token_ids_buffer_ptr,
+ top_logprobs_buffer_ptr=top_logprobs_buffer_ptr,
+ token_count=token_count,
+ kv_cache_size=kv_cache_size,
+ TOPK=topk,
+ MAX_TOPK=max_topk,
+ BLOCK=block,
+ )
diff --git a/lightllm/common/basemodel/triton_kernel/routing_capture.py b/lightllm/common/basemodel/triton_kernel/routing_capture.py
new file mode 100644
index 0000000000..d0fa822058
--- /dev/null
+++ b/lightllm/common/basemodel/triton_kernel/routing_capture.py
@@ -0,0 +1,74 @@
+import torch
+import triton
+import triton.language as tl
+
+
+@triton.jit
+def _scatter_routing_topk_to_cpu(
+ topk_ids,
+ mem_indexes,
+ routing_buffer_ptr,
+ total_size,
+ moe_layer_index: tl.constexpr,
+ layer_topk_size: tl.constexpr,
+ topk: tl.constexpr,
+ dtype_id: tl.constexpr,
+ BLOCK: tl.constexpr,
+):
+ pid = tl.program_id(0)
+ offsets = pid * BLOCK + tl.arange(0, BLOCK)
+ mask = offsets < total_size
+
+ token_offsets = offsets // topk
+ topk_offsets = offsets - token_offsets * topk
+ mem_index = tl.load(mem_indexes + token_offsets, mask=mask, other=-1).to(tl.int64)
+ data = tl.load(topk_ids + offsets, mask=mask, other=0)
+
+ dst_offsets = mem_index * layer_topk_size + moe_layer_index * topk + topk_offsets
+ if dtype_id == 1:
+ dst_ptr = tl.load(routing_buffer_ptr).to(tl.pointer_type(tl.uint8))
+ tl.store(dst_ptr + dst_offsets, data.to(tl.uint8), mask=mask)
+ else:
+ dst_ptr = tl.load(routing_buffer_ptr).to(tl.pointer_type(tl.int16))
+ tl.store(dst_ptr + dst_offsets, data.to(tl.int16), mask=mask)
+
+
+def scatter_routing_topk_to_cpu(
+ topk_ids: torch.Tensor,
+ mem_indexes: torch.Tensor,
+ routing_buffer_ptr: torch.Tensor,
+ moe_layer_index: int,
+ num_moe_layers: int,
+ topk: int,
+ dtype_id: int,
+):
+ assert topk_ids.is_cuda
+ assert mem_indexes.is_cuda
+ assert mem_indexes.is_contiguous()
+ assert routing_buffer_ptr.is_cuda
+ assert routing_buffer_ptr.dtype == torch.uint64
+ assert routing_buffer_ptr.numel() == 1
+ assert topk_ids.dim() == 2
+ assert topk_ids.shape[1] == topk
+ assert topk_ids.is_contiguous()
+ assert 0 <= moe_layer_index < num_moe_layers
+
+ num_tokens = topk_ids.shape[0]
+ layer_topk_size = num_moe_layers * topk
+ total_size = num_tokens * topk
+ if total_size == 0:
+ return
+
+ BLOCK = 1024
+ grid = (triton.cdiv(total_size, BLOCK),)
+ _scatter_routing_topk_to_cpu[grid](
+ topk_ids=topk_ids,
+ mem_indexes=mem_indexes,
+ routing_buffer_ptr=routing_buffer_ptr,
+ total_size=total_size,
+ moe_layer_index=moe_layer_index,
+ layer_topk_size=layer_topk_size,
+ topk=topk,
+ dtype_id=dtype_id,
+ BLOCK=BLOCK,
+ )
diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.4.0/NVIDIA_H200/grouped_matmul:v1/{K=2048,N=768,expert_num=128,mul_routed_weight=false,out_dtype=torch.bfloat16,topk_num=8,use_fp8_w8a8=false}_NVIDIA_H200.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.4.0/NVIDIA_H200/grouped_matmul:v1/{K=2048,N=768,expert_num=128,mul_routed_weight=false,out_dtype=torch.bfloat16,topk_num=8,use_fp8_w8a8=false}_NVIDIA_H200.json
new file mode 100644
index 0000000000..c75c871c72
--- /dev/null
+++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.4.0/NVIDIA_H200/grouped_matmul:v1/{K=2048,N=768,expert_num=128,mul_routed_weight=false,out_dtype=torch.bfloat16,topk_num=8,use_fp8_w8a8=false}_NVIDIA_H200.json
@@ -0,0 +1,110 @@
+{
+ "1": {
+ "BLOCK_SIZE_K": 128,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 16,
+ "NEED_TRANS": false,
+ "num_stages": 5,
+ "num_warps": 4
+ },
+ "100": {
+ "BLOCK_SIZE_K": 128,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 1,
+ "NEED_TRANS": false,
+ "num_stages": 2,
+ "num_warps": 8
+ },
+ "1024": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 128,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 1,
+ "NEED_TRANS": false,
+ "num_stages": 3,
+ "num_warps": 4
+ },
+ "128": {
+ "BLOCK_SIZE_K": 128,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 64,
+ "GROUP_SIZE_M": 64,
+ "NEED_TRANS": false,
+ "num_stages": 2,
+ "num_warps": 4
+ },
+ "16": {
+ "BLOCK_SIZE_K": 128,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 64,
+ "GROUP_SIZE_M": 32,
+ "NEED_TRANS": false,
+ "num_stages": 3,
+ "num_warps": 4
+ },
+ "2048": {
+ "BLOCK_SIZE_K": 32,
+ "BLOCK_SIZE_M": 128,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 1,
+ "NEED_TRANS": false,
+ "num_stages": 4,
+ "num_warps": 4
+ },
+ "256": {
+ "BLOCK_SIZE_K": 128,
+ "BLOCK_SIZE_M": 32,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 1,
+ "NEED_TRANS": false,
+ "num_stages": 3,
+ "num_warps": 4
+ },
+ "32": {
+ "BLOCK_SIZE_K": 128,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 1,
+ "NEED_TRANS": false,
+ "num_stages": 2,
+ "num_warps": 4
+ },
+ "4096": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 128,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 1,
+ "NEED_TRANS": false,
+ "num_stages": 3,
+ "num_warps": 4
+ },
+ "64": {
+ "BLOCK_SIZE_K": 128,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 64,
+ "GROUP_SIZE_M": 32,
+ "NEED_TRANS": false,
+ "num_stages": 2,
+ "num_warps": 4
+ },
+ "8": {
+ "BLOCK_SIZE_K": 128,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 64,
+ "GROUP_SIZE_M": 1,
+ "NEED_TRANS": false,
+ "num_stages": 3,
+ "num_warps": 4
+ },
+ "8448": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 128,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 1,
+ "NEED_TRANS": false,
+ "num_stages": 3,
+ "num_warps": 4
+ }
+}
\ No newline at end of file
diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.4.0/NVIDIA_H200/grouped_matmul:v1/{K=384,N=2048,expert_num=128,mul_routed_weight=true,out_dtype=torch.bfloat16,topk_num=1,use_fp8_w8a8=false}_NVIDIA_H200.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.4.0/NVIDIA_H200/grouped_matmul:v1/{K=384,N=2048,expert_num=128,mul_routed_weight=true,out_dtype=torch.bfloat16,topk_num=1,use_fp8_w8a8=false}_NVIDIA_H200.json
new file mode 100644
index 0000000000..14026090e6
--- /dev/null
+++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.4.0/NVIDIA_H200/grouped_matmul:v1/{K=384,N=2048,expert_num=128,mul_routed_weight=true,out_dtype=torch.bfloat16,topk_num=1,use_fp8_w8a8=false}_NVIDIA_H200.json
@@ -0,0 +1,110 @@
+{
+ "1024": {
+ "BLOCK_SIZE_K": 128,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 1,
+ "NEED_TRANS": false,
+ "num_stages": 2,
+ "num_warps": 4
+ },
+ "128": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 64,
+ "GROUP_SIZE_M": 1,
+ "NEED_TRANS": false,
+ "num_stages": 4,
+ "num_warps": 4
+ },
+ "16384": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 128,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 16,
+ "NEED_TRANS": false,
+ "num_stages": 3,
+ "num_warps": 4
+ },
+ "2048": {
+ "BLOCK_SIZE_K": 128,
+ "BLOCK_SIZE_M": 32,
+ "BLOCK_SIZE_N": 64,
+ "GROUP_SIZE_M": 16,
+ "NEED_TRANS": false,
+ "num_stages": 2,
+ "num_warps": 4
+ },
+ "256": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 64,
+ "GROUP_SIZE_M": 1,
+ "NEED_TRANS": false,
+ "num_stages": 4,
+ "num_warps": 4
+ },
+ "32768": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 128,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 16,
+ "NEED_TRANS": false,
+ "num_stages": 3,
+ "num_warps": 4
+ },
+ "512": {
+ "BLOCK_SIZE_K": 128,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 16,
+ "NEED_TRANS": false,
+ "num_stages": 2,
+ "num_warps": 4
+ },
+ "64": {
+ "BLOCK_SIZE_K": 128,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 64,
+ "GROUP_SIZE_M": 16,
+ "NEED_TRANS": false,
+ "num_stages": 2,
+ "num_warps": 4
+ },
+ "67584": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 128,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 16,
+ "NEED_TRANS": false,
+ "num_stages": 3,
+ "num_warps": 4
+ },
+ "8": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 1,
+ "NEED_TRANS": false,
+ "num_stages": 3,
+ "num_warps": 8
+ },
+ "800": {
+ "BLOCK_SIZE_K": 128,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 64,
+ "GROUP_SIZE_M": 1,
+ "NEED_TRANS": false,
+ "num_stages": 2,
+ "num_warps": 4
+ },
+ "8192": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 64,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 16,
+ "NEED_TRANS": false,
+ "num_stages": 3,
+ "num_warps": 4
+ }
+}
\ No newline at end of file
diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.4.0/NVIDIA_H200/moe_sum_reduce:v1/{hidden_dim=2048,out_dtype=torch.bfloat16,topk_num=8}_NVIDIA_H200.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.4.0/NVIDIA_H200/moe_sum_reduce:v1/{hidden_dim=2048,out_dtype=torch.bfloat16,topk_num=8}_NVIDIA_H200.json
new file mode 100644
index 0000000000..939c939523
--- /dev/null
+++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.4.0/NVIDIA_H200/moe_sum_reduce:v1/{hidden_dim=2048,out_dtype=torch.bfloat16,topk_num=8}_NVIDIA_H200.json
@@ -0,0 +1,74 @@
+{
+ "1": {
+ "BLOCK_DIM": 128,
+ "BLOCK_M": 2,
+ "NUM_STAGE": 2,
+ "num_warps": 4
+ },
+ "100": {
+ "BLOCK_DIM": 512,
+ "BLOCK_M": 1,
+ "NUM_STAGE": 1,
+ "num_warps": 8
+ },
+ "1024": {
+ "BLOCK_DIM": 1024,
+ "BLOCK_M": 1,
+ "NUM_STAGE": 1,
+ "num_warps": 4
+ },
+ "128": {
+ "BLOCK_DIM": 1024,
+ "BLOCK_M": 1,
+ "NUM_STAGE": 1,
+ "num_warps": 16
+ },
+ "16": {
+ "BLOCK_DIM": 128,
+ "BLOCK_M": 1,
+ "NUM_STAGE": 2,
+ "num_warps": 4
+ },
+ "2048": {
+ "BLOCK_DIM": 256,
+ "BLOCK_M": 1,
+ "NUM_STAGE": 1,
+ "num_warps": 2
+ },
+ "256": {
+ "BLOCK_DIM": 1024,
+ "BLOCK_M": 1,
+ "NUM_STAGE": 1,
+ "num_warps": 8
+ },
+ "32": {
+ "BLOCK_DIM": 256,
+ "BLOCK_M": 1,
+ "NUM_STAGE": 1,
+ "num_warps": 2
+ },
+ "4096": {
+ "BLOCK_DIM": 512,
+ "BLOCK_M": 1,
+ "NUM_STAGE": 1,
+ "num_warps": 4
+ },
+ "64": {
+ "BLOCK_DIM": 1024,
+ "BLOCK_M": 1,
+ "NUM_STAGE": 1,
+ "num_warps": 8
+ },
+ "8": {
+ "BLOCK_DIM": 128,
+ "BLOCK_M": 1,
+ "NUM_STAGE": 1,
+ "num_warps": 8
+ },
+ "8448": {
+ "BLOCK_DIM": 1024,
+ "BLOCK_M": 1,
+ "NUM_STAGE": 1,
+ "num_warps": 4
+ }
+}
\ No newline at end of file
diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.4.0/NVIDIA_H200/silu_and_mul_fwd:v1/{N=384,out_dtype=torch.bfloat16}_NVIDIA_H200.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.4.0/NVIDIA_H200/silu_and_mul_fwd:v1/{N=384,out_dtype=torch.bfloat16}_NVIDIA_H200.json
new file mode 100644
index 0000000000..13ba4ba8e5
--- /dev/null
+++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.4.0/NVIDIA_H200/silu_and_mul_fwd:v1/{N=384,out_dtype=torch.bfloat16}_NVIDIA_H200.json
@@ -0,0 +1,74 @@
+{
+ "1024": {
+ "BLOCK_M": 8,
+ "BLOCK_N": 128,
+ "NUM_STAGES": 4,
+ "num_warps": 1
+ },
+ "128": {
+ "BLOCK_M": 1,
+ "BLOCK_N": 256,
+ "NUM_STAGES": 1,
+ "num_warps": 4
+ },
+ "16384": {
+ "BLOCK_M": 8,
+ "BLOCK_N": 256,
+ "NUM_STAGES": 4,
+ "num_warps": 1
+ },
+ "2048": {
+ "BLOCK_M": 8,
+ "BLOCK_N": 256,
+ "NUM_STAGES": 4,
+ "num_warps": 1
+ },
+ "256": {
+ "BLOCK_M": 1,
+ "BLOCK_N": 256,
+ "NUM_STAGES": 2,
+ "num_warps": 4
+ },
+ "32768": {
+ "BLOCK_M": 32,
+ "BLOCK_N": 128,
+ "NUM_STAGES": 4,
+ "num_warps": 1
+ },
+ "512": {
+ "BLOCK_M": 1,
+ "BLOCK_N": 256,
+ "NUM_STAGES": 1,
+ "num_warps": 8
+ },
+ "64": {
+ "BLOCK_M": 1,
+ "BLOCK_N": 128,
+ "NUM_STAGES": 1,
+ "num_warps": 1
+ },
+ "67584": {
+ "BLOCK_M": 64,
+ "BLOCK_N": 128,
+ "NUM_STAGES": 4,
+ "num_warps": 1
+ },
+ "8": {
+ "BLOCK_M": 1,
+ "BLOCK_N": 128,
+ "NUM_STAGES": 4,
+ "num_warps": 4
+ },
+ "800": {
+ "BLOCK_M": 1,
+ "BLOCK_N": 256,
+ "NUM_STAGES": 1,
+ "num_warps": 4
+ },
+ "8192": {
+ "BLOCK_M": 8,
+ "BLOCK_N": 256,
+ "NUM_STAGES": 4,
+ "num_warps": 1
+ }
+}
\ No newline at end of file
diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H100_80GB_HBM3/grouped_matmul:v1/{K=192,N=2048,expert_num=128,mul_routed_weight=true,out_dtype=torch.bfloat16,topk_num=1,use_fp8_w8a8=false}_NVIDIA_H100_80GB_HBM3.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H100_80GB_HBM3/grouped_matmul:v1/{K=192,N=2048,expert_num=128,mul_routed_weight=true,out_dtype=torch.bfloat16,topk_num=1,use_fp8_w8a8=false}_NVIDIA_H100_80GB_HBM3.json
new file mode 100644
index 0000000000..ee316f610b
--- /dev/null
+++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H100_80GB_HBM3/grouped_matmul:v1/{K=192,N=2048,expert_num=128,mul_routed_weight=true,out_dtype=torch.bfloat16,topk_num=1,use_fp8_w8a8=false}_NVIDIA_H100_80GB_HBM3.json
@@ -0,0 +1,110 @@
+{
+ "1024": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 64,
+ "GROUP_SIZE_M": 1,
+ "NEED_TRANS": false,
+ "num_stages": 3,
+ "num_warps": 8
+ },
+ "128": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 64,
+ "GROUP_SIZE_M": 64,
+ "NEED_TRANS": false,
+ "num_stages": 3,
+ "num_warps": 4
+ },
+ "16384": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 64,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 32,
+ "NEED_TRANS": false,
+ "num_stages": 2,
+ "num_warps": 4
+ },
+ "2048": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 32,
+ "BLOCK_SIZE_N": 64,
+ "GROUP_SIZE_M": 1,
+ "NEED_TRANS": false,
+ "num_stages": 3,
+ "num_warps": 4
+ },
+ "256": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 64,
+ "GROUP_SIZE_M": 1,
+ "NEED_TRANS": false,
+ "num_stages": 2,
+ "num_warps": 4
+ },
+ "32768": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 64,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 32,
+ "NEED_TRANS": false,
+ "num_stages": 2,
+ "num_warps": 4
+ },
+ "512": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 1,
+ "NEED_TRANS": false,
+ "num_stages": 4,
+ "num_warps": 4
+ },
+ "64": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 16,
+ "NEED_TRANS": false,
+ "num_stages": 2,
+ "num_warps": 4
+ },
+ "67584": {
+ "BLOCK_SIZE_K": 32,
+ "BLOCK_SIZE_M": 64,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 16,
+ "NEED_TRANS": false,
+ "num_stages": 4,
+ "num_warps": 4
+ },
+ "8": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 1,
+ "NEED_TRANS": false,
+ "num_stages": 3,
+ "num_warps": 4
+ },
+ "800": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 64,
+ "GROUP_SIZE_M": 1,
+ "NEED_TRANS": false,
+ "num_stages": 3,
+ "num_warps": 8
+ },
+ "8192": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 64,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 16,
+ "NEED_TRANS": false,
+ "num_stages": 2,
+ "num_warps": 4
+ }
+}
\ No newline at end of file
diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H100_80GB_HBM3/grouped_matmul:v1/{K=192,N=2048,expert_num=128,mul_routed_weight=true,out_dtype=torch.float16,topk_num=1,use_fp8_w8a8=false}_NVIDIA_H100_80GB_HBM3.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H100_80GB_HBM3/grouped_matmul:v1/{K=192,N=2048,expert_num=128,mul_routed_weight=true,out_dtype=torch.float16,topk_num=1,use_fp8_w8a8=false}_NVIDIA_H100_80GB_HBM3.json
new file mode 100644
index 0000000000..e027701092
--- /dev/null
+++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H100_80GB_HBM3/grouped_matmul:v1/{K=192,N=2048,expert_num=128,mul_routed_weight=true,out_dtype=torch.float16,topk_num=1,use_fp8_w8a8=false}_NVIDIA_H100_80GB_HBM3.json
@@ -0,0 +1,110 @@
+{
+ "1024": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 1,
+ "NEED_TRANS": false,
+ "num_stages": 2,
+ "num_warps": 4
+ },
+ "128": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 1,
+ "NEED_TRANS": false,
+ "num_stages": 4,
+ "num_warps": 4
+ },
+ "16384": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 64,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 32,
+ "NEED_TRANS": false,
+ "num_stages": 2,
+ "num_warps": 4
+ },
+ "2048": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 32,
+ "BLOCK_SIZE_N": 64,
+ "GROUP_SIZE_M": 1,
+ "NEED_TRANS": false,
+ "num_stages": 3,
+ "num_warps": 4
+ },
+ "256": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 32,
+ "GROUP_SIZE_M": 1,
+ "NEED_TRANS": false,
+ "num_stages": 3,
+ "num_warps": 4
+ },
+ "32768": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 64,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 32,
+ "NEED_TRANS": false,
+ "num_stages": 2,
+ "num_warps": 4
+ },
+ "512": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 16,
+ "NEED_TRANS": false,
+ "num_stages": 2,
+ "num_warps": 4
+ },
+ "64": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 64,
+ "GROUP_SIZE_M": 32,
+ "NEED_TRANS": false,
+ "num_stages": 3,
+ "num_warps": 4
+ },
+ "67584": {
+ "BLOCK_SIZE_K": 32,
+ "BLOCK_SIZE_M": 64,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 16,
+ "NEED_TRANS": false,
+ "num_stages": 4,
+ "num_warps": 4
+ },
+ "8": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 1,
+ "NEED_TRANS": false,
+ "num_stages": 3,
+ "num_warps": 4
+ },
+ "800": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 1,
+ "NEED_TRANS": false,
+ "num_stages": 2,
+ "num_warps": 4
+ },
+ "8192": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 64,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 16,
+ "NEED_TRANS": false,
+ "num_stages": 2,
+ "num_warps": 4
+ }
+}
\ No newline at end of file
diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H100_80GB_HBM3/grouped_matmul:v1/{K=192,N=2048,expert_num=128,mul_routed_weight=true,out_dtype=torch.float16,topk_num=1,use_fp8_w8a8=true}_NVIDIA_H100_80GB_HBM3.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H100_80GB_HBM3/grouped_matmul:v1/{K=192,N=2048,expert_num=128,mul_routed_weight=true,out_dtype=torch.float16,topk_num=1,use_fp8_w8a8=true}_NVIDIA_H100_80GB_HBM3.json
new file mode 100644
index 0000000000..ddda23d257
--- /dev/null
+++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H100_80GB_HBM3/grouped_matmul:v1/{K=192,N=2048,expert_num=128,mul_routed_weight=true,out_dtype=torch.float16,topk_num=1,use_fp8_w8a8=true}_NVIDIA_H100_80GB_HBM3.json
@@ -0,0 +1,110 @@
+{
+ "1024": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 64,
+ "NEED_TRANS": true,
+ "num_stages": 4,
+ "num_warps": 4
+ },
+ "128": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 1,
+ "NEED_TRANS": true,
+ "num_stages": 2,
+ "num_warps": 4
+ },
+ "16384": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 64,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 64,
+ "NEED_TRANS": false,
+ "num_stages": 4,
+ "num_warps": 4
+ },
+ "2048": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 32,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 64,
+ "NEED_TRANS": true,
+ "num_stages": 4,
+ "num_warps": 4
+ },
+ "256": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 16,
+ "NEED_TRANS": true,
+ "num_stages": 2,
+ "num_warps": 4
+ },
+ "32768": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 64,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 32,
+ "NEED_TRANS": false,
+ "num_stages": 4,
+ "num_warps": 4
+ },
+ "512": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 32,
+ "NEED_TRANS": false,
+ "num_stages": 4,
+ "num_warps": 4
+ },
+ "64": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 1,
+ "NEED_TRANS": true,
+ "num_stages": 2,
+ "num_warps": 4
+ },
+ "67584": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 64,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 16,
+ "NEED_TRANS": false,
+ "num_stages": 4,
+ "num_warps": 4
+ },
+ "8": {
+ "BLOCK_SIZE_K": 128,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 1,
+ "NEED_TRANS": true,
+ "num_stages": 2,
+ "num_warps": 4
+ },
+ "800": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 32,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 64,
+ "NEED_TRANS": true,
+ "num_stages": 4,
+ "num_warps": 4
+ },
+ "8192": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 64,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 64,
+ "NEED_TRANS": false,
+ "num_stages": 4,
+ "num_warps": 4
+ }
+}
\ No newline at end of file
diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H100_80GB_HBM3/grouped_matmul:v1/{K=2048,N=384,expert_num=128,mul_routed_weight=false,out_dtype=torch.bfloat16,topk_num=8,use_fp8_w8a8=false}_NVIDIA_H100_80GB_HBM3.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H100_80GB_HBM3/grouped_matmul:v1/{K=2048,N=384,expert_num=128,mul_routed_weight=false,out_dtype=torch.bfloat16,topk_num=8,use_fp8_w8a8=false}_NVIDIA_H100_80GB_HBM3.json
new file mode 100644
index 0000000000..560ca6c09d
--- /dev/null
+++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H100_80GB_HBM3/grouped_matmul:v1/{K=2048,N=384,expert_num=128,mul_routed_weight=false,out_dtype=torch.bfloat16,topk_num=8,use_fp8_w8a8=false}_NVIDIA_H100_80GB_HBM3.json
@@ -0,0 +1,110 @@
+{
+ "1": {
+ "BLOCK_SIZE_K": 128,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 64,
+ "GROUP_SIZE_M": 1,
+ "NEED_TRANS": false,
+ "num_stages": 5,
+ "num_warps": 4
+ },
+ "100": {
+ "BLOCK_SIZE_K": 128,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 16,
+ "NEED_TRANS": false,
+ "num_stages": 2,
+ "num_warps": 4
+ },
+ "1024": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 64,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 32,
+ "NEED_TRANS": false,
+ "num_stages": 3,
+ "num_warps": 4
+ },
+ "128": {
+ "BLOCK_SIZE_K": 128,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 16,
+ "NEED_TRANS": false,
+ "num_stages": 2,
+ "num_warps": 8
+ },
+ "16": {
+ "BLOCK_SIZE_K": 128,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 1,
+ "NEED_TRANS": false,
+ "num_stages": 3,
+ "num_warps": 4
+ },
+ "2048": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 64,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 16,
+ "NEED_TRANS": false,
+ "num_stages": 2,
+ "num_warps": 4
+ },
+ "256": {
+ "BLOCK_SIZE_K": 128,
+ "BLOCK_SIZE_M": 32,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 64,
+ "NEED_TRANS": false,
+ "num_stages": 2,
+ "num_warps": 4
+ },
+ "32": {
+ "BLOCK_SIZE_K": 128,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 64,
+ "GROUP_SIZE_M": 1,
+ "NEED_TRANS": false,
+ "num_stages": 2,
+ "num_warps": 4
+ },
+ "4096": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 128,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 32,
+ "NEED_TRANS": false,
+ "num_stages": 3,
+ "num_warps": 4
+ },
+ "64": {
+ "BLOCK_SIZE_K": 128,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 32,
+ "NEED_TRANS": false,
+ "num_stages": 2,
+ "num_warps": 4
+ },
+ "8": {
+ "BLOCK_SIZE_K": 128,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 16,
+ "NEED_TRANS": false,
+ "num_stages": 3,
+ "num_warps": 4
+ },
+ "8448": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 128,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 1,
+ "NEED_TRANS": false,
+ "num_stages": 3,
+ "num_warps": 8
+ }
+}
\ No newline at end of file
diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H100_80GB_HBM3/grouped_matmul:v1/{K=2048,N=384,expert_num=128,mul_routed_weight=false,out_dtype=torch.float16,topk_num=8,use_fp8_w8a8=false}_NVIDIA_H100_80GB_HBM3.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H100_80GB_HBM3/grouped_matmul:v1/{K=2048,N=384,expert_num=128,mul_routed_weight=false,out_dtype=torch.float16,topk_num=8,use_fp8_w8a8=false}_NVIDIA_H100_80GB_HBM3.json
new file mode 100644
index 0000000000..0713de7996
--- /dev/null
+++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H100_80GB_HBM3/grouped_matmul:v1/{K=2048,N=384,expert_num=128,mul_routed_weight=false,out_dtype=torch.float16,topk_num=8,use_fp8_w8a8=false}_NVIDIA_H100_80GB_HBM3.json
@@ -0,0 +1,110 @@
+{
+ "1": {
+ "BLOCK_SIZE_K": 128,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 64,
+ "GROUP_SIZE_M": 1,
+ "NEED_TRANS": false,
+ "num_stages": 3,
+ "num_warps": 4
+ },
+ "100": {
+ "BLOCK_SIZE_K": 128,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 1,
+ "NEED_TRANS": false,
+ "num_stages": 2,
+ "num_warps": 4
+ },
+ "1024": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 64,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 32,
+ "NEED_TRANS": false,
+ "num_stages": 3,
+ "num_warps": 4
+ },
+ "128": {
+ "BLOCK_SIZE_K": 128,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 16,
+ "NEED_TRANS": false,
+ "num_stages": 2,
+ "num_warps": 8
+ },
+ "16": {
+ "BLOCK_SIZE_K": 128,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 1,
+ "NEED_TRANS": false,
+ "num_stages": 3,
+ "num_warps": 4
+ },
+ "2048": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 64,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 16,
+ "NEED_TRANS": false,
+ "num_stages": 2,
+ "num_warps": 4
+ },
+ "256": {
+ "BLOCK_SIZE_K": 128,
+ "BLOCK_SIZE_M": 32,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 16,
+ "NEED_TRANS": false,
+ "num_stages": 2,
+ "num_warps": 4
+ },
+ "32": {
+ "BLOCK_SIZE_K": 128,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 64,
+ "GROUP_SIZE_M": 16,
+ "NEED_TRANS": false,
+ "num_stages": 3,
+ "num_warps": 4
+ },
+ "4096": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 128,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 32,
+ "NEED_TRANS": false,
+ "num_stages": 3,
+ "num_warps": 4
+ },
+ "64": {
+ "BLOCK_SIZE_K": 128,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 32,
+ "NEED_TRANS": false,
+ "num_stages": 2,
+ "num_warps": 4
+ },
+ "8": {
+ "BLOCK_SIZE_K": 128,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 16,
+ "NEED_TRANS": false,
+ "num_stages": 3,
+ "num_warps": 4
+ },
+ "8448": {
+ "BLOCK_SIZE_K": 64,
+ "BLOCK_SIZE_M": 128,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 1,
+ "NEED_TRANS": false,
+ "num_stages": 3,
+ "num_warps": 8
+ }
+}
\ No newline at end of file
diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H100_80GB_HBM3/grouped_matmul:v1/{K=2048,N=384,expert_num=128,mul_routed_weight=false,out_dtype=torch.float16,topk_num=8,use_fp8_w8a8=true}_NVIDIA_H100_80GB_HBM3.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H100_80GB_HBM3/grouped_matmul:v1/{K=2048,N=384,expert_num=128,mul_routed_weight=false,out_dtype=torch.float16,topk_num=8,use_fp8_w8a8=true}_NVIDIA_H100_80GB_HBM3.json
new file mode 100644
index 0000000000..e950ff0954
--- /dev/null
+++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H100_80GB_HBM3/grouped_matmul:v1/{K=2048,N=384,expert_num=128,mul_routed_weight=false,out_dtype=torch.float16,topk_num=8,use_fp8_w8a8=true}_NVIDIA_H100_80GB_HBM3.json
@@ -0,0 +1,110 @@
+{
+ "1": {
+ "BLOCK_SIZE_K": 128,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 64,
+ "GROUP_SIZE_M": 16,
+ "NEED_TRANS": true,
+ "num_stages": 3,
+ "num_warps": 4
+ },
+ "100": {
+ "BLOCK_SIZE_K": 128,
+ "BLOCK_SIZE_M": 32,
+ "BLOCK_SIZE_N": 64,
+ "GROUP_SIZE_M": 16,
+ "NEED_TRANS": true,
+ "num_stages": 3,
+ "num_warps": 4
+ },
+ "1024": {
+ "BLOCK_SIZE_K": 128,
+ "BLOCK_SIZE_M": 64,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 32,
+ "NEED_TRANS": false,
+ "num_stages": 3,
+ "num_warps": 4
+ },
+ "128": {
+ "BLOCK_SIZE_K": 128,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 64,
+ "GROUP_SIZE_M": 64,
+ "NEED_TRANS": true,
+ "num_stages": 3,
+ "num_warps": 4
+ },
+ "16": {
+ "BLOCK_SIZE_K": 128,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 64,
+ "NEED_TRANS": true,
+ "num_stages": 3,
+ "num_warps": 4
+ },
+ "2048": {
+ "BLOCK_SIZE_K": 128,
+ "BLOCK_SIZE_M": 128,
+ "BLOCK_SIZE_N": 64,
+ "GROUP_SIZE_M": 32,
+ "NEED_TRANS": true,
+ "num_stages": 3,
+ "num_warps": 4
+ },
+ "256": {
+ "BLOCK_SIZE_K": 128,
+ "BLOCK_SIZE_M": 32,
+ "BLOCK_SIZE_N": 64,
+ "GROUP_SIZE_M": 64,
+ "NEED_TRANS": true,
+ "num_stages": 3,
+ "num_warps": 4
+ },
+ "32": {
+ "BLOCK_SIZE_K": 128,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 64,
+ "NEED_TRANS": true,
+ "num_stages": 3,
+ "num_warps": 4
+ },
+ "4096": {
+ "BLOCK_SIZE_K": 128,
+ "BLOCK_SIZE_M": 128,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 16,
+ "NEED_TRANS": false,
+ "num_stages": 3,
+ "num_warps": 8
+ },
+ "64": {
+ "BLOCK_SIZE_K": 128,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 64,
+ "GROUP_SIZE_M": 1,
+ "NEED_TRANS": true,
+ "num_stages": 3,
+ "num_warps": 4
+ },
+ "8": {
+ "BLOCK_SIZE_K": 128,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 32,
+ "NEED_TRANS": true,
+ "num_stages": 4,
+ "num_warps": 4
+ },
+ "8448": {
+ "BLOCK_SIZE_K": 128,
+ "BLOCK_SIZE_M": 128,
+ "BLOCK_SIZE_N": 128,
+ "GROUP_SIZE_M": 1,
+ "NEED_TRANS": false,
+ "num_stages": 3,
+ "num_warps": 8
+ }
+}
\ No newline at end of file
diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H100_80GB_HBM3/moe_align_fused:v1/{topk_num=8}_NVIDIA_H100_80GB_HBM3.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H100_80GB_HBM3/moe_align_fused:v1/{topk_num=8}_NVIDIA_H100_80GB_HBM3.json
new file mode 100644
index 0000000000..7f479b8382
--- /dev/null
+++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H100_80GB_HBM3/moe_align_fused:v1/{topk_num=8}_NVIDIA_H100_80GB_HBM3.json
@@ -0,0 +1,50 @@
+{
+ "1": {
+ "BLOCK_SIZE": 256,
+ "num_warps": 2
+ },
+ "100": {
+ "BLOCK_SIZE": 128,
+ "num_warps": 8
+ },
+ "1024": {
+ "BLOCK_SIZE": 128,
+ "num_warps": 2
+ },
+ "128": {
+ "BLOCK_SIZE": 128,
+ "num_warps": 8
+ },
+ "16": {
+ "BLOCK_SIZE": 256,
+ "num_warps": 4
+ },
+ "2048": {
+ "BLOCK_SIZE": 256,
+ "num_warps": 8
+ },
+ "256": {
+ "BLOCK_SIZE": 128,
+ "num_warps": 8
+ },
+ "32": {
+ "BLOCK_SIZE": 128,
+ "num_warps": 8
+ },
+ "4096": {
+ "BLOCK_SIZE": 128,
+ "num_warps": 1
+ },
+ "64": {
+ "BLOCK_SIZE": 128,
+ "num_warps": 8
+ },
+ "8": {
+ "BLOCK_SIZE": 256,
+ "num_warps": 2
+ },
+ "8448": {
+ "BLOCK_SIZE": 256,
+ "num_warps": 4
+ }
+}
\ No newline at end of file
diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H100_80GB_HBM3/moe_sum_reduce:v1/{hidden_dim=2048,out_dtype=torch.bfloat16,topk_num=8}_NVIDIA_H100_80GB_HBM3.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H100_80GB_HBM3/moe_sum_reduce:v1/{hidden_dim=2048,out_dtype=torch.bfloat16,topk_num=8}_NVIDIA_H100_80GB_HBM3.json
new file mode 100644
index 0000000000..b3051c6584
--- /dev/null
+++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H100_80GB_HBM3/moe_sum_reduce:v1/{hidden_dim=2048,out_dtype=torch.bfloat16,topk_num=8}_NVIDIA_H100_80GB_HBM3.json
@@ -0,0 +1,74 @@
+{
+ "1": {
+ "BLOCK_DIM": 128,
+ "BLOCK_M": 1,
+ "NUM_STAGE": 2,
+ "num_warps": 8
+ },
+ "100": {
+ "BLOCK_DIM": 512,
+ "BLOCK_M": 1,
+ "NUM_STAGE": 1,
+ "num_warps": 8
+ },
+ "1024": {
+ "BLOCK_DIM": 1024,
+ "BLOCK_M": 2,
+ "NUM_STAGE": 4,
+ "num_warps": 4
+ },
+ "128": {
+ "BLOCK_DIM": 512,
+ "BLOCK_M": 1,
+ "NUM_STAGE": 1,
+ "num_warps": 8
+ },
+ "16": {
+ "BLOCK_DIM": 128,
+ "BLOCK_M": 1,
+ "NUM_STAGE": 2,
+ "num_warps": 4
+ },
+ "2048": {
+ "BLOCK_DIM": 512,
+ "BLOCK_M": 1,
+ "NUM_STAGE": 1,
+ "num_warps": 4
+ },
+ "256": {
+ "BLOCK_DIM": 1024,
+ "BLOCK_M": 1,
+ "NUM_STAGE": 4,
+ "num_warps": 2
+ },
+ "32": {
+ "BLOCK_DIM": 256,
+ "BLOCK_M": 1,
+ "NUM_STAGE": 1,
+ "num_warps": 2
+ },
+ "4096": {
+ "BLOCK_DIM": 256,
+ "BLOCK_M": 1,
+ "NUM_STAGE": 1,
+ "num_warps": 1
+ },
+ "64": {
+ "BLOCK_DIM": 512,
+ "BLOCK_M": 1,
+ "NUM_STAGE": 1,
+ "num_warps": 4
+ },
+ "8": {
+ "BLOCK_DIM": 256,
+ "BLOCK_M": 1,
+ "NUM_STAGE": 1,
+ "num_warps": 2
+ },
+ "8448": {
+ "BLOCK_DIM": 256,
+ "BLOCK_M": 1,
+ "NUM_STAGE": 1,
+ "num_warps": 2
+ }
+}
\ No newline at end of file
diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H100_80GB_HBM3/moe_sum_reduce:v1/{hidden_dim=2048,out_dtype=torch.float16,topk_num=8}_NVIDIA_H100_80GB_HBM3.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H100_80GB_HBM3/moe_sum_reduce:v1/{hidden_dim=2048,out_dtype=torch.float16,topk_num=8}_NVIDIA_H100_80GB_HBM3.json
new file mode 100644
index 0000000000..fdb3212216
--- /dev/null
+++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H100_80GB_HBM3/moe_sum_reduce:v1/{hidden_dim=2048,out_dtype=torch.float16,topk_num=8}_NVIDIA_H100_80GB_HBM3.json
@@ -0,0 +1,74 @@
+{
+ "1": {
+ "BLOCK_DIM": 256,
+ "BLOCK_M": 1,
+ "NUM_STAGE": 4,
+ "num_warps": 8
+ },
+ "100": {
+ "BLOCK_DIM": 1024,
+ "BLOCK_M": 1,
+ "NUM_STAGE": 1,
+ "num_warps": 8
+ },
+ "1024": {
+ "BLOCK_DIM": 512,
+ "BLOCK_M": 1,
+ "NUM_STAGE": 1,
+ "num_warps": 2
+ },
+ "128": {
+ "BLOCK_DIM": 1024,
+ "BLOCK_M": 1,
+ "NUM_STAGE": 1,
+ "num_warps": 8
+ },
+ "16": {
+ "BLOCK_DIM": 128,
+ "BLOCK_M": 1,
+ "NUM_STAGE": 4,
+ "num_warps": 4
+ },
+ "2048": {
+ "BLOCK_DIM": 1024,
+ "BLOCK_M": 1,
+ "NUM_STAGE": 1,
+ "num_warps": 8
+ },
+ "256": {
+ "BLOCK_DIM": 1024,
+ "BLOCK_M": 1,
+ "NUM_STAGE": 1,
+ "num_warps": 8
+ },
+ "32": {
+ "BLOCK_DIM": 256,
+ "BLOCK_M": 1,
+ "NUM_STAGE": 2,
+ "num_warps": 8
+ },
+ "4096": {
+ "BLOCK_DIM": 256,
+ "BLOCK_M": 1,
+ "NUM_STAGE": 1,
+ "num_warps": 2
+ },
+ "64": {
+ "BLOCK_DIM": 1024,
+ "BLOCK_M": 1,
+ "NUM_STAGE": 1,
+ "num_warps": 8
+ },
+ "8": {
+ "BLOCK_DIM": 128,
+ "BLOCK_M": 1,
+ "NUM_STAGE": 1,
+ "num_warps": 4
+ },
+ "8448": {
+ "BLOCK_DIM": 1024,
+ "BLOCK_M": 1,
+ "NUM_STAGE": 1,
+ "num_warps": 8
+ }
+}
\ No newline at end of file
diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H100_80GB_HBM3/silu_and_mul_fwd:v1/{N=192,out_dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H100_80GB_HBM3/silu_and_mul_fwd:v1/{N=192,out_dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json
new file mode 100644
index 0000000000..a94e669353
--- /dev/null
+++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H100_80GB_HBM3/silu_and_mul_fwd:v1/{N=192,out_dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json
@@ -0,0 +1,74 @@
+{
+ "1024": {
+ "BLOCK_M": 1,
+ "BLOCK_N": 256,
+ "NUM_STAGES": 1,
+ "num_warps": 4
+ },
+ "128": {
+ "BLOCK_M": 1,
+ "BLOCK_N": 256,
+ "NUM_STAGES": 4,
+ "num_warps": 1
+ },
+ "16384": {
+ "BLOCK_M": 8,
+ "BLOCK_N": 256,
+ "NUM_STAGES": 4,
+ "num_warps": 1
+ },
+ "2048": {
+ "BLOCK_M": 1,
+ "BLOCK_N": 256,
+ "NUM_STAGES": 1,
+ "num_warps": 4
+ },
+ "256": {
+ "BLOCK_M": 1,
+ "BLOCK_N": 256,
+ "NUM_STAGES": 1,
+ "num_warps": 4
+ },
+ "32768": {
+ "BLOCK_M": 1,
+ "BLOCK_N": 256,
+ "NUM_STAGES": 4,
+ "num_warps": 1
+ },
+ "512": {
+ "BLOCK_M": 1,
+ "BLOCK_N": 256,
+ "NUM_STAGES": 1,
+ "num_warps": 4
+ },
+ "64": {
+ "BLOCK_M": 1,
+ "BLOCK_N": 64,
+ "NUM_STAGES": 1,
+ "num_warps": 4
+ },
+ "67584": {
+ "BLOCK_M": 1,
+ "BLOCK_N": 256,
+ "NUM_STAGES": 4,
+ "num_warps": 1
+ },
+ "8": {
+ "BLOCK_M": 1,
+ "BLOCK_N": 64,
+ "NUM_STAGES": 1,
+ "num_warps": 4
+ },
+ "800": {
+ "BLOCK_M": 1,
+ "BLOCK_N": 256,
+ "NUM_STAGES": 1,
+ "num_warps": 4
+ },
+ "8192": {
+ "BLOCK_M": 8,
+ "BLOCK_N": 256,
+ "NUM_STAGES": 4,
+ "num_warps": 1
+ }
+}
\ No newline at end of file
diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H100_80GB_HBM3/silu_and_mul_fwd:v1/{N=192,out_dtype=torch.float16}_NVIDIA_H100_80GB_HBM3.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H100_80GB_HBM3/silu_and_mul_fwd:v1/{N=192,out_dtype=torch.float16}_NVIDIA_H100_80GB_HBM3.json
new file mode 100644
index 0000000000..441421fd5d
--- /dev/null
+++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H100_80GB_HBM3/silu_and_mul_fwd:v1/{N=192,out_dtype=torch.float16}_NVIDIA_H100_80GB_HBM3.json
@@ -0,0 +1,74 @@
+{
+ "1024": {
+ "BLOCK_M": 1,
+ "BLOCK_N": 256,
+ "NUM_STAGES": 1,
+ "num_warps": 1
+ },
+ "128": {
+ "BLOCK_M": 1,
+ "BLOCK_N": 64,
+ "NUM_STAGES": 1,
+ "num_warps": 1
+ },
+ "16384": {
+ "BLOCK_M": 8,
+ "BLOCK_N": 256,
+ "NUM_STAGES": 4,
+ "num_warps": 1
+ },
+ "2048": {
+ "BLOCK_M": 1,
+ "BLOCK_N": 256,
+ "NUM_STAGES": 1,
+ "num_warps": 1
+ },
+ "256": {
+ "BLOCK_M": 1,
+ "BLOCK_N": 256,
+ "NUM_STAGES": 1,
+ "num_warps": 1
+ },
+ "32768": {
+ "BLOCK_M": 1,
+ "BLOCK_N": 256,
+ "NUM_STAGES": 1,
+ "num_warps": 1
+ },
+ "512": {
+ "BLOCK_M": 1,
+ "BLOCK_N": 256,
+ "NUM_STAGES": 1,
+ "num_warps": 1
+ },
+ "64": {
+ "BLOCK_M": 1,
+ "BLOCK_N": 64,
+ "NUM_STAGES": 1,
+ "num_warps": 1
+ },
+ "67584": {
+ "BLOCK_M": 1,
+ "BLOCK_N": 256,
+ "NUM_STAGES": 4,
+ "num_warps": 1
+ },
+ "8": {
+ "BLOCK_M": 1,
+ "BLOCK_N": 64,
+ "NUM_STAGES": 1,
+ "num_warps": 1
+ },
+ "800": {
+ "BLOCK_M": 1,
+ "BLOCK_N": 256,
+ "NUM_STAGES": 1,
+ "num_warps": 4
+ },
+ "8192": {
+ "BLOCK_M": 8,
+ "BLOCK_N": 256,
+ "NUM_STAGES": 4,
+ "num_warps": 1
+ }
+}
\ No newline at end of file
diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.1/NVIDIA_H100_80GB_HBM3/gemma_rmsnorm_forward:v1/{N=2048,weight_dtype=torch.bfloat16,x_dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.1/NVIDIA_H100_80GB_HBM3/gemma_rmsnorm_forward:v1/{N=2048,weight_dtype=torch.bfloat16,x_dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json
new file mode 100644
index 0000000000..864d1d3f18
--- /dev/null
+++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.1/NVIDIA_H100_80GB_HBM3/gemma_rmsnorm_forward:v1/{N=2048,weight_dtype=torch.bfloat16,x_dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json
@@ -0,0 +1,7 @@
+{
+ "2048": {
+ "BLOCK_SIZE": 4096,
+ "num_stages": 1,
+ "num_warps": 8
+ }
+}
\ No newline at end of file
diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.1/NVIDIA_H100_80GB_HBM3/gemma_rmsnorm_forward:v1/{N=256,weight_dtype=torch.bfloat16,x_dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.1/NVIDIA_H100_80GB_HBM3/gemma_rmsnorm_forward:v1/{N=256,weight_dtype=torch.bfloat16,x_dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json
new file mode 100644
index 0000000000..bcf56e01f7
--- /dev/null
+++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.1/NVIDIA_H100_80GB_HBM3/gemma_rmsnorm_forward:v1/{N=256,weight_dtype=torch.bfloat16,x_dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json
@@ -0,0 +1,7 @@
+{
+ "256": {
+ "BLOCK_SIZE": 128,
+ "num_stages": 1,
+ "num_warps": 1
+ }
+}
\ No newline at end of file
diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.1/NVIDIA_H100_80GB_HBM3/gemma_rmsnorm_forward:v1/{N=3072,weight_dtype=torch.bfloat16,x_dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.1/NVIDIA_H100_80GB_HBM3/gemma_rmsnorm_forward:v1/{N=3072,weight_dtype=torch.bfloat16,x_dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json
new file mode 100644
index 0000000000..ba1dc8a75d
--- /dev/null
+++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.1/NVIDIA_H100_80GB_HBM3/gemma_rmsnorm_forward:v1/{N=3072,weight_dtype=torch.bfloat16,x_dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json
@@ -0,0 +1,7 @@
+{
+ "3072": {
+ "BLOCK_SIZE": 2048,
+ "num_stages": 1,
+ "num_warps": 8
+ }
+}
\ No newline at end of file
diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.1/NVIDIA_H100_80GB_HBM3/gemma_rmsnorm_forward:v1/{N=5120,weight_dtype=torch.bfloat16,x_dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.1/NVIDIA_H100_80GB_HBM3/gemma_rmsnorm_forward:v1/{N=5120,weight_dtype=torch.bfloat16,x_dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json
new file mode 100644
index 0000000000..6f109e1c6e
--- /dev/null
+++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.1/NVIDIA_H100_80GB_HBM3/gemma_rmsnorm_forward:v1/{N=5120,weight_dtype=torch.bfloat16,x_dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json
@@ -0,0 +1,7 @@
+{
+ "5120": {
+ "BLOCK_SIZE": 32768,
+ "num_stages": 1,
+ "num_warps": 8
+ }
+}
\ No newline at end of file
diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.1/NVIDIA_H200/gemma_rmsnorm_forward:v1/{N=2048,weight_dtype=torch.bfloat16,x_dtype=torch.bfloat16}_NVIDIA_H200.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.1/NVIDIA_H200/gemma_rmsnorm_forward:v1/{N=2048,weight_dtype=torch.bfloat16,x_dtype=torch.bfloat16}_NVIDIA_H200.json
new file mode 100644
index 0000000000..198a196dfb
--- /dev/null
+++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.1/NVIDIA_H200/gemma_rmsnorm_forward:v1/{N=2048,weight_dtype=torch.bfloat16,x_dtype=torch.bfloat16}_NVIDIA_H200.json
@@ -0,0 +1,7 @@
+{
+ "2048": {
+ "BLOCK_SIZE": 1024,
+ "num_stages": 1,
+ "num_warps": 4
+ }
+}
\ No newline at end of file
diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.1/NVIDIA_H200/gemma_rmsnorm_forward:v1/{N=256,weight_dtype=torch.bfloat16,x_dtype=torch.bfloat16}_NVIDIA_H200.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.1/NVIDIA_H200/gemma_rmsnorm_forward:v1/{N=256,weight_dtype=torch.bfloat16,x_dtype=torch.bfloat16}_NVIDIA_H200.json
new file mode 100644
index 0000000000..537c7a90eb
--- /dev/null
+++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.1/NVIDIA_H200/gemma_rmsnorm_forward:v1/{N=256,weight_dtype=torch.bfloat16,x_dtype=torch.bfloat16}_NVIDIA_H200.json
@@ -0,0 +1,7 @@
+{
+ "256": {
+ "BLOCK_SIZE": 512,
+ "num_stages": 1,
+ "num_warps": 1
+ }
+}
\ No newline at end of file
diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.1/NVIDIA_H200/gemma_rmsnorm_forward:v1/{N=4096,weight_dtype=torch.bfloat16,x_dtype=torch.bfloat16}_NVIDIA_H200.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.1/NVIDIA_H200/gemma_rmsnorm_forward:v1/{N=4096,weight_dtype=torch.bfloat16,x_dtype=torch.bfloat16}_NVIDIA_H200.json
new file mode 100644
index 0000000000..9a6dcb6fbf
--- /dev/null
+++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.1/NVIDIA_H200/gemma_rmsnorm_forward:v1/{N=4096,weight_dtype=torch.bfloat16,x_dtype=torch.bfloat16}_NVIDIA_H200.json
@@ -0,0 +1,7 @@
+{
+ "4096": {
+ "BLOCK_SIZE": 1024,
+ "num_stages": 1,
+ "num_warps": 8
+ }
+}
\ No newline at end of file
diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.1/NVIDIA_H200/gemma_rmsnorm_forward:v1/{N=5120,weight_dtype=torch.bfloat16,x_dtype=torch.bfloat16}_NVIDIA_H200.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.1/NVIDIA_H200/gemma_rmsnorm_forward:v1/{N=5120,weight_dtype=torch.bfloat16,x_dtype=torch.bfloat16}_NVIDIA_H200.json
new file mode 100644
index 0000000000..df501847ec
--- /dev/null
+++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.1/NVIDIA_H200/gemma_rmsnorm_forward:v1/{N=5120,weight_dtype=torch.bfloat16,x_dtype=torch.bfloat16}_NVIDIA_H200.json
@@ -0,0 +1,7 @@
+{
+ "5120": {
+ "BLOCK_SIZE": 1024,
+ "num_stages": 1,
+ "num_warps": 8
+ }
+}
\ No newline at end of file
diff --git a/lightllm/common/triton_utils/autotuner.py b/lightllm/common/triton_utils/autotuner.py
index c62a2572ff..4cc6453d12 100644
--- a/lightllm/common/triton_utils/autotuner.py
+++ b/lightllm/common/triton_utils/autotuner.py
@@ -11,7 +11,7 @@
from frozendict import frozendict
from lightllm.utils.device_utils import get_current_device_name
from lightllm.utils.log_utils import init_logger
-from typing import Callable, Optional, Union, List
+from typing import Callable, List
from lightllm.utils.envs_utils import get_triton_autotune_level
from lightllm.common.kernel_config import KernelConfigs
from lightllm.utils.dist_utils import get_global_world_size, get_global_rank, get_current_rank_in_node
@@ -106,14 +106,6 @@ def __init__(
self.configs_gen_func = configs_gen_func
self.kernel_name = kernel_name
- self.cache_dir = os.path.join(
- Path(__file__).parent,
- "autotune_kernel_configs",
- get_triton_version(),
- get_current_device_name(),
- self.kernel_name,
- )
- os.makedirs(self.cache_dir, exist_ok=True)
self.fn = fn
self.static_key_func = static_key_func
self.run_key_func = run_key_func
@@ -209,6 +201,25 @@ def __call__(self, *args, **kwargs):
return self.fn(*args, **kwargs)
+ @property
+ def cache_dir(self) -> str:
+ if not hasattr(self, "_cache_dir"):
+ device_name = get_current_device_name()
+ if device_name is None:
+ raise RuntimeError(
+ f"Autotuner for kernel {self.kernel_name} requires a visible CUDA/MUSA device "
+ f"to resolve its cache directory, but torch.cuda.is_available() is False."
+ )
+ self._cache_dir = os.path.join(
+ Path(__file__).parent,
+ "autotune_kernel_configs",
+ get_triton_version(),
+ device_name,
+ self.kernel_name,
+ )
+ os.makedirs(self._cache_dir, exist_ok=True)
+ return self._cache_dir
+
def _try_load_cache(self, static_key):
if static_key in self.cached_configs:
return False
diff --git a/lightllm/models/deepseek2/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek2/layer_infer/transformer_layer_infer.py
index be819c94a0..dae79cc8a6 100644
--- a/lightllm/models/deepseek2/layer_infer/transformer_layer_infer.py
+++ b/lightllm/models/deepseek2/layer_infer/transformer_layer_infer.py
@@ -232,6 +232,7 @@ def _moe_ffn_tp(
use_grouped_topk=self.n_group,
topk_group=self.topk_group,
num_expert_group=self.n_group,
+ infer_state=infer_state,
)
if self.n_shared_experts is not None and layer_weight.num_fused_shared_experts == 0:
@@ -259,6 +260,7 @@ def _moe_ffn_edp(
topk_group=self.topk_group,
num_expert_group=self.n_group,
is_prefill=infer_state.is_prefill,
+ infer_state=infer_state,
)
if self.n_shared_experts is not None:
diff --git a/lightllm/models/gemma4/layer_infer/post_layer_infer.py b/lightllm/models/gemma4/layer_infer/post_layer_infer.py
index 22bcf0508d..b736a2d6c1 100644
--- a/lightllm/models/gemma4/layer_infer/post_layer_infer.py
+++ b/lightllm/models/gemma4/layer_infer/post_layer_infer.py
@@ -17,4 +17,6 @@ def token_forward(self, input_embdings, infer_state, layer_weight):
if self.final_logit_softcapping is not None and self.final_logit_softcapping > 0:
cap = self.final_logit_softcapping
logits = torch.tanh(logits / cap) * cap
+ if infer_state.prompt_logics is not None:
+ infer_state.prompt_logics = torch.tanh(infer_state.prompt_logics / cap) * cap
return logits
diff --git a/lightllm/models/gemma4/layer_infer/transformer_layer_infer.py b/lightllm/models/gemma4/layer_infer/transformer_layer_infer.py
index 015b526fbc..2f0c01dbf6 100644
--- a/lightllm/models/gemma4/layer_infer/transformer_layer_infer.py
+++ b/lightllm/models/gemma4/layer_infer/transformer_layer_infer.py
@@ -300,6 +300,7 @@ def _ffn_moe(self, input, router_logits, infer_state: InferStateInfo, layer_weig
topk_group=None,
num_expert_group=None,
is_prefill=infer_state.is_prefill,
+ infer_state=infer_state,
)
moe_out = self._tpsp_reduce(input=moe_out, infer_state=infer_state)
return moe_out
diff --git a/lightllm/models/gpt_oss/layer_infer/transformer_layer_infer.py b/lightllm/models/gpt_oss/layer_infer/transformer_layer_infer.py
index b27ea8fd2d..490d2dc4c5 100644
--- a/lightllm/models/gpt_oss/layer_infer/transformer_layer_infer.py
+++ b/lightllm/models/gpt_oss/layer_infer/transformer_layer_infer.py
@@ -52,6 +52,7 @@ def _ffn(self, input, infer_state, layer_weight: GptOssTransformerLayerWeight) -
use_grouped_topk=False,
topk_group=None,
num_expert_group=None,
+ infer_state=infer_state,
)
hidden_states = hidden_states.view(num_tokens, hidden_dim)
return self._tpsp_reduce(input=hidden_states, infer_state=infer_state)
diff --git a/lightllm/models/llama/layer_infer/post_layer_infer.py b/lightllm/models/llama/layer_infer/post_layer_infer.py
index 50dc0109e2..bb6e4f3735 100644
--- a/lightllm/models/llama/layer_infer/post_layer_infer.py
+++ b/lightllm/models/llama/layer_infer/post_layer_infer.py
@@ -39,18 +39,23 @@ def _slice_get_last_input(self, input_embdings: torch.Tensor, infer_state: Llama
last_input[:, :] = input_embdings[last_index, :]
return last_input, select_token_num
- if infer_state.is_prefill and not infer_state.return_all_prompt_logics:
+ if infer_state.is_prefill:
+ # logits 始终只取每个请求最后一个位置的 hidden state,用于正常采样。
batch_size = infer_state.batch_size
last_input = self.alloc_tensor((batch_size, embed_dim_), dtype=input_embdings.dtype)
last_index = (
torch.cumsum(infer_state.b_seq_len - infer_state.b_ready_cache_len, dim=0, dtype=torch.long) - 1
)
last_input[:, :] = input_embdings[last_index, :]
- return last_input, batch_size
- if infer_state.is_prefill and infer_state.return_all_prompt_logics:
- total_tokens = infer_state.total_token_num
- return input_embdings, total_tokens
+ # 在开启 return_all_prompt_logics 模式时,额外保存整个 prefill 阶段
+ # 每一个 token 位置对应的 hidden state,用于后续输出 prompt logprobs。
+ # input_embdings 本身已经是本次新增的 token(不含已缓存前缀),
+ # 仅在 chunked prefill 的 padding 场景下会多出行,padding 部分会在
+ # basemodel._create_unpad_prefill_model_output 中按实际 token 数量裁剪掉。
+ if infer_state.return_all_prompt_logics:
+ infer_state.prompt_logics = input_embdings
+ return last_input, batch_size
if not infer_state.is_prefill:
batch_size = infer_state.batch_size
@@ -62,17 +67,41 @@ def _token_forward(
self, input_embdings: torch.Tensor, infer_state: LlamaInferStateInfo, layer_weight: LlamaPreAndPostLayerWeight
):
last_input, token_num = self._slice_get_last_input(input_embdings, infer_state)
- input_embdings_dtype = input_embdings.dtype
input_embdings = None
- last_input = self._norm(last_input, infer_state, layer_weight)
- last_input = last_input.permute(1, 0).view(-1, token_num)
- logic_batch = layer_weight.lm_head_weight_(input=last_input, alloc_func=self.alloc_tensor)
- last_input = None
+
+ # 正常采样使用的 logits,始终只对应每个请求最后一个位置。
+ ans_logics = self._lm_head_and_gather(last_input, token_num, layer_weight, infer_state)
+ # 在 return_all_prompt_logics 模式下,prompt_logics 保存的是完整 prefill
+ # 的 hidden state,需要在 norm/lm_head 之前取出来,避免被 input_embdings 置空。
+ prompt_logics_hiddens = infer_state.prompt_logics
+ infer_state.prompt_logics = None
+ # 在 return_all_prompt_logics 模式下,额外计算整个 prefill 阶段所有位置的 logits,
+ # 存入返回的 prompt_logics 中,原来的 ans_logics 仅保留最后一个位置的 logits。
+ if prompt_logics_hiddens is not None:
+ prompt_token_num = prompt_logics_hiddens.shape[0]
+ infer_state.prompt_logics = self._lm_head_and_gather(
+ prompt_logics_hiddens, prompt_token_num, layer_weight, infer_state
+ )
+
+ return ans_logics
+
+ def _lm_head_and_gather(
+ self,
+ hidden: torch.Tensor,
+ token_num: int,
+ layer_weight: LlamaPreAndPostLayerWeight,
+ infer_state: LlamaInferStateInfo,
+ ) -> torch.Tensor:
+ normed = self._norm(hidden, infer_state, layer_weight)
+ normed = normed.permute(1, 0).view(-1, token_num)
+ logic_batch = layer_weight.lm_head_weight_(input=normed, alloc_func=self.alloc_tensor)
+ normed = None
+
vocab_size = layer_weight.lm_head_weight_.vocab_size
if self.tp_world_size_ == 1:
gather_data = logic_batch
else:
- gather_data = self.alloc_tensor((vocab_size, token_num), dtype=input_embdings_dtype)
+ gather_data = self.alloc_tensor((vocab_size, token_num), dtype=hidden.dtype)
split_indexes = np.linspace(0, vocab_size, self.tp_world_size_ + 1, dtype=np.int64)
all_gather(
[gather_data[split_indexes[i] : split_indexes[i + 1], :] for i in range(self.tp_world_size_)],
@@ -81,10 +110,8 @@ def _token_forward(
async_op=False,
)
logic_batch = None
- ans_logics = self.alloc_tensor(
- (token_num, vocab_size),
- dtype=torch.float32,
- )
+
+ ans_logics = self.alloc_tensor((token_num, vocab_size), dtype=torch.float32)
ans_logics[:, :] = gather_data.permute(1, 0)
gather_data = None
return ans_logics
diff --git a/lightllm/models/mixtral/layer_infer/_custom_ops.py b/lightllm/models/mixtral/layer_infer/_custom_ops.py
deleted file mode 100644
index b0e27ac1de..0000000000
--- a/lightllm/models/mixtral/layer_infer/_custom_ops.py
+++ /dev/null
@@ -1,46 +0,0 @@
-import functools
-import json
-import os
-from typing import Any, Dict, Optional, Tuple
-
-import torch
-import triton
-import triton.language as tl
-from lightllm.utils.log_utils import init_logger
-
-logger = init_logger(__name__)
-
-# Pytorch version
-# Triton version in progress
-def topk_softmax(
- topk_weights,
- topk_ids,
- token_expert_indicies,
- gating_output,
- topk=2,
-):
- scores = torch.softmax(gating_output, dim=-1)
- topk_weights, topk_ids = torch.topk(scores, k=topk, dim=-1, sorted=False)
- return topk_weights, topk_ids
-
-
-def fused_topk(
- hidden_states: torch.Tensor,
- gating_output: torch.Tensor,
- topk: int,
- renormalize: bool,
- alloc_tensor_func=torch.empty,
-):
- assert hidden_states.shape[0] == gating_output.shape[0], "Number of tokens mismatch"
-
- M, _ = hidden_states.shape
-
- topk_weights = alloc_tensor_func((M, topk), dtype=torch.float32, device=hidden_states.device)
- topk_ids = alloc_tensor_func((M, topk), dtype=torch.int32, device=hidden_states.device)
- token_expert_indicies = alloc_tensor_func((M, topk), dtype=torch.int32, device=hidden_states.device)
- topk_weights, topk_ids = topk_softmax(topk_weights, topk_ids, token_expert_indicies, gating_output.float(), topk)
- del token_expert_indicies # Not used. Will be used in the future.
-
- if renormalize:
- topk_weights = topk_weights / topk_weights.sum(dim=-1, keepdim=True)
- return topk_weights, topk_ids
diff --git a/lightllm/models/mixtral/layer_infer/transformer_layer_infer.py b/lightllm/models/mixtral/layer_infer/transformer_layer_infer.py
index 0cf651598a..8134dc266d 100644
--- a/lightllm/models/mixtral/layer_infer/transformer_layer_infer.py
+++ b/lightllm/models/mixtral/layer_infer/transformer_layer_infer.py
@@ -1,9 +1,6 @@
-import os
import torch
-import torch.nn.functional as F
from lightllm.common.basemodel.infer_struct import InferStateInfo
from lightllm.models.llama.layer_infer.transformer_layer_infer import LlamaTransformerLayerInfer
-from lightllm.models.mixtral.layer_infer._custom_ops import fused_topk
from lightllm.models.mixtral.layer_weights.transformer_layer_weight import MixtralTransformerLayerWeight
@@ -21,25 +18,14 @@ def _ffn(self, input, infer_state: InferStateInfo, layer_weight: MixtralTransfor
num_tokens, hidden_dim = hidden_states.shape
router_logits = layer_weight.moe_gate.mm(hidden_states)
- topk_weights, topk_ids = fused_topk(
- hidden_states=hidden_states,
- gating_output=router_logits,
- topk=self.num_experts_per_tok,
+ layer_weight.experts.experts(
+ hidden_states,
+ router_logits=router_logits,
+ top_k=self.num_experts_per_tok,
renormalize=self.renormalize,
- alloc_tensor_func=self.alloc_tensor,
+ use_grouped_topk=False,
+ topk_group=None,
+ num_expert_group=None,
+ infer_state=infer_state,
)
- from lightllm.common.fused_moe.grouped_fused_moe import fused_experts_impl
-
- ffn2_out = fused_experts_impl(
- hidden_states=hidden_states,
- w1=layer_weight.experts.w1[0],
- w2=layer_weight.experts.w2[0],
- topk_weights=topk_weights,
- topk_ids=topk_ids,
- inplace=True,
- use_fp8_w8a8=False,
- w1_scale=None,
- w2_scale=None,
- alloc_tensor_func=self.alloc_tensor,
- )
- return self._tpsp_reduce(input=ffn2_out, infer_state=infer_state)
+ return hidden_states.view(num_tokens, hidden_dim)
diff --git a/lightllm/models/qwen2_vl/model.py b/lightllm/models/qwen2_vl/model.py
index 237c4ad897..c94135573b 100644
--- a/lightllm/models/qwen2_vl/model.py
+++ b/lightllm/models/qwen2_vl/model.py
@@ -12,6 +12,7 @@
from .vision_process import smart_resize
from lightllm.models.qwen2.model import Qwen2TpPartModel
import os
+from typing import Union, List
# Warp of the origal tokenizer
class QWen2VLTokenizer(BaseMultiModalTokenizer):
@@ -52,9 +53,13 @@ def get_image_token_length(self, img: ImageItem):
def get_audio_token_length(self, audio: AudioItem):
raise NotImplementedError
- def encode(self, prompt, multimodal_params: MultimodalParams = None, **kwargs):
-
- origin_ids = self.tokenizer.encode(prompt)
+ def encode(self, prompt: Union[str, List[int]], multimodal_params: MultimodalParams = None, **kwargs):
+ if isinstance(prompt, str):
+ origin_ids = self.tokenizer.encode(prompt)
+ elif isinstance(prompt, list):
+ origin_ids = prompt
+ else:
+ raise ValueError(f"Unsupported prompt type: {type(prompt)}")
#
->
origin_ids = [token for token in origin_ids if token != self.image_token_id]
diff --git a/lightllm/models/qwen3_5_moe/layer_infer/__init__.py b/lightllm/models/qwen3_5_moe/layer_infer/__init__.py
new file mode 100644
index 0000000000..e69de29bb2
diff --git a/lightllm/models/qwen3_5_moe/layer_weights/transformer_layer_weight.py b/lightllm/models/qwen3_5_moe/layer_weights/transformer_layer_weight.py
index fe4b1883bd..7a91ed47e3 100644
--- a/lightllm/models/qwen3_5_moe/layer_weights/transformer_layer_weight.py
+++ b/lightllm/models/qwen3_5_moe/layer_weights/transformer_layer_weight.py
@@ -1,40 +1,42 @@
-import torch
from lightllm.models.qwen3_5.layer_weights.transformer_layer_weight import Qwen35TransformerLayerWeight
class Qwen35MOETransformerLayerWeight(Qwen35TransformerLayerWeight):
def load_hf_weights(self, weights):
- moe_intermediate_size = self.network_config_["moe_intermediate_size"]
- split_fused_expert_weights(weights, self.layer_num_, moe_intermediate_size)
+ split_fused_expert_weights(weights, self.layer_num_, self.network_config_["moe_intermediate_size"])
return super().load_hf_weights(weights)
def split_fused_expert_weights(weights: dict, layer_num: int, moe_intermediate_size: int):
+ """将 HF 打包的 fused MoE expert 权重拆成按 expert 索引的独立权重。
+
+ 部分 checkpoint(如 Qwen3.5-MoE)把所有 expert 的 gate_up / down 压成
+ ``mlp.experts.{gate_up,down}_proj`` 的打包张量(首维为 expert 数)。
+ 本函数只处理 ``model.layers.{layer_num}`` 下的这类 key:弹出打包权重,
+ 再写入 ``mlp.experts.{expert_idx}.{gate,up,down}_proj.weight``,供后续
+ 按 expert 加载。``gate_up_proj`` 会按 ``moe_intermediate_size`` 沿
+ intermediate 维切成 gate / up。
+ """
layer_prefix = f"model.layers.{layer_num}."
keys = list(weights.keys())
- num_experts = 0
for k in keys:
if not k.startswith(layer_prefix):
continue
if "mlp.experts.gate_up_proj" in k:
- fused_weight = weights.pop(k) # [num_experts, 2*inter_size, hidden_size]
- num_experts = fused_weight.shape[0]
-
+ fused_weight = weights.pop(k)
prefix = k.rsplit(".gate_up_proj", 1)[0]
gate_weight = fused_weight[:, :moe_intermediate_size, :]
up_weight = fused_weight[:, moe_intermediate_size:, :]
- for expert_idx in range(num_experts):
+ for expert_idx in range(fused_weight.shape[0]):
weights[f"{prefix}.{expert_idx}.gate_proj.weight"] = gate_weight[expert_idx]
weights[f"{prefix}.{expert_idx}.up_proj.weight"] = up_weight[expert_idx]
elif "mlp.experts.down_proj" in k:
- down_weight = weights.pop(k) # [num_experts, hidden_size, inter_size]
- num_experts = down_weight.shape[0]
-
+ down_weight = weights.pop(k)
prefix = k.rsplit(".down_proj", 1)[0]
- for expert_idx in range(num_experts):
+ for expert_idx in range(down_weight.shape[0]):
weights[f"{prefix}.{expert_idx}.down_proj.weight"] = down_weight[expert_idx]
diff --git a/lightllm/models/qwen3_moe/layer_infer/transformer_layer_infer.py b/lightllm/models/qwen3_moe/layer_infer/transformer_layer_infer.py
index 8879aa2d27..7edfd5a6f9 100644
--- a/lightllm/models/qwen3_moe/layer_infer/transformer_layer_infer.py
+++ b/lightllm/models/qwen3_moe/layer_infer/transformer_layer_infer.py
@@ -86,6 +86,7 @@ def _moe_ffn_tp(
use_grouped_topk=False,
topk_group=None,
num_expert_group=None,
+ infer_state=infer_state,
)
return hidden_states.view(num_tokens, hidden_dim)
@@ -105,6 +106,7 @@ def _moe_ffn_edp(
topk_group=None,
num_expert_group=None,
is_prefill=infer_state.is_prefill,
+ infer_state=infer_state,
)
ep_output = ep_output.view(token_num, hidden_dim)
diff --git a/lightllm/models/qwen3_moe/layer_weights/transformer_layer_weight.py b/lightllm/models/qwen3_moe/layer_weights/transformer_layer_weight.py
index e525cb2d20..38ba17b244 100644
--- a/lightllm/models/qwen3_moe/layer_weights/transformer_layer_weight.py
+++ b/lightllm/models/qwen3_moe/layer_weights/transformer_layer_weight.py
@@ -1,4 +1,3 @@
-import os
from lightllm.models.qwen3.layer_weights.transformer_layer_weight import Qwen3TransformerLayerWeight
from lightllm.common.basemodel.layer_weights.meta_weights import ROWMMWeight, FusedMoeWeight, QKVROWNMMWeight
@@ -14,6 +13,15 @@ def __init__(self, layer_num, data_type, network_config, quant_cfg=None):
super().__init__(layer_num, data_type, network_config, quant_cfg)
return
+ def load_hf_weights(self, weights):
+ if self.is_moe:
+ split_fused_expert_weights(
+ weights,
+ self.layer_num_,
+ self.network_config_["moe_intermediate_size"],
+ )
+ return super().load_hf_weights(weights)
+
def _init_weight_names(self):
self._q_weight_name = f"model.layers.{self.layer_num_}.self_attn.q_proj.weight"
self._q_norm_name = f"model.layers.{self.layer_num_}.self_attn.q_norm.weight"
@@ -79,3 +87,50 @@ def _init_qkv(self):
bias_names=[self._q_bias_name, self._k_bias_name, self._v_bias_name],
quant_method=self.get_quant_method("qkv_proj"),
)
+
+
+def split_fused_expert_weights(weights: dict, layer_num: int, moe_intermediate_size: int):
+ """将 HF 打包的 fused MoE expert 权重拆成按 expert 索引的独立权重。
+
+ 部分 checkpoint(如 Qwen3-MoE)把所有 expert 的 gate/up/down 压成
+ ``mlp.experts.{gate_up,gate,up,down}_proj`` 的打包张量
+ (首维为 expert 数)。本函数只处理 ``model.layers.{layer_num}`` 下的这类
+ key:弹出打包权重,再写入
+ ``mlp.experts.{expert_idx}.{gate,up,down}_proj.weight``,供后续按 expert
+ 加载。若存在 fused ``gate_up_proj``,还会按 ``moe_intermediate_size``
+ 沿 intermediate 维切成 gate / up。
+ """
+ layer_prefix = f"model.layers.{layer_num}."
+ keys = list(weights.keys())
+
+ for k in keys:
+ if not k.startswith(layer_prefix):
+ continue
+
+ if "mlp.experts.gate_up_proj" in k:
+ fused_weight = weights.pop(k)
+ prefix = k.rsplit(".gate_up_proj", 1)[0]
+ gate_weight = fused_weight[:, :moe_intermediate_size, :]
+ up_weight = fused_weight[:, moe_intermediate_size:, :]
+
+ for expert_idx in range(fused_weight.shape[0]):
+ weights[f"{prefix}.{expert_idx}.gate_proj.weight"] = gate_weight[expert_idx]
+ weights[f"{prefix}.{expert_idx}.up_proj.weight"] = up_weight[expert_idx]
+
+ elif "mlp.experts.gate_proj" in k:
+ gate_weight = weights.pop(k)
+ prefix = k.rsplit(".gate_proj", 1)[0]
+ for expert_idx in range(gate_weight.shape[0]):
+ weights[f"{prefix}.{expert_idx}.gate_proj.weight"] = gate_weight[expert_idx]
+
+ elif "mlp.experts.up_proj" in k:
+ up_weight = weights.pop(k)
+ prefix = k.rsplit(".up_proj", 1)[0]
+ for expert_idx in range(up_weight.shape[0]):
+ weights[f"{prefix}.{expert_idx}.up_proj.weight"] = up_weight[expert_idx]
+
+ elif "mlp.experts.down_proj" in k:
+ down_weight = weights.pop(k)
+ prefix = k.rsplit(".down_proj", 1)[0]
+ for expert_idx in range(down_weight.shape[0]):
+ weights[f"{prefix}.{expert_idx}.down_proj.weight"] = down_weight[expert_idx]
diff --git a/lightllm/models/qwen3_vl_moe/layer_weights/transformers_layer_weight.py b/lightllm/models/qwen3_vl_moe/layer_weights/transformers_layer_weight.py
index 83c05ba264..9af6d93869 100644
--- a/lightllm/models/qwen3_vl_moe/layer_weights/transformers_layer_weight.py
+++ b/lightllm/models/qwen3_vl_moe/layer_weights/transformers_layer_weight.py
@@ -1,42 +1,35 @@
-import os
from lightllm.models.qwen3_moe.layer_weights.transformer_layer_weight import Qwen3MOETransformerLayerWeight
class Qwen3VLMOETransformerLayerWeight(Qwen3MOETransformerLayerWeight):
- def __init__(self, layer_num, data_type, network_config, quant_cfg=None):
- super().__init__(layer_num, data_type, network_config, quant_cfg)
-
def load_hf_weights(self, weights):
+ self._align_fused_expert_weight_layout(weights)
+ return super().load_hf_weights(weights)
+
+ def _align_fused_expert_weight_layout(self, weights: dict) -> None:
+ """将 Qwen3-VL-MoE 的 fused expert 权重布局对齐到基类期望格式。
+
+ Qwen3-VL-MoE 官方 safetensor 中,expert 权重以 packed 3D 张量存储,但维度
+ 顺序与 Qwen3-MoE 文本模型 / ``nn.Linear`` 权重布局不同:
+
+ - safetensor(VL):
+ - ``gate_up_proj``: ``[E, H, 2I]``
+ - ``down_proj``: ``[E, I, H]``
+ - 基类 ``split_fused_expert_weights`` 期望(与 Linear 一致):
+ - ``gate_up_proj``: ``[E, 2I, H]``
+ - ``down_proj``: ``[E, H, I]``
+
+ 本函数仅对当前层存在的 fused key 做 ``transpose(1, 2)`` 布局转换,
+ 不负责按 expert 拆分;拆分仍复用基类逻辑。若对应 key 不在本批
+ ``weights`` 中(分片加载),则跳过。
+ """
moe_prefix = f"model.layers.{self.layer_num_}.mlp.experts"
gate_up_name = f"{moe_prefix}.gate_up_proj"
down_name = f"{moe_prefix}.down_proj"
if gate_up_name in weights:
- gate_up = weights[gate_up_name] # [E, H, 2I]
- E, H, twoI = gate_up.shape
- assert twoI % 2 == 0, f"gate_up_proj last dim must be even, but got {twoI}"
- I_dim = twoI // 2
-
- if down_name in weights:
- down = weights[down_name] # [E, I, H]
- else:
- down = None
-
- for e in range(E):
- gate_up_e = gate_up[e]
- gate_e = gate_up_e[:, :I_dim].transpose(0, 1).contiguous()
- up_e = gate_up_e[:, I_dim:].transpose(0, 1).contiguous()
-
- gate_key = f"{moe_prefix}.{e}.gate_proj.weight"
- up_key = f"{moe_prefix}.{e}.up_proj.weight"
- weights[gate_key] = gate_e
- weights[up_key] = up_e
-
- if down is not None:
- down_key = f"{moe_prefix}.{e}.down_proj.weight"
- weights[down_key] = down[e].transpose(0, 1).contiguous()
-
- del weights[gate_up_name]
- if down_name in weights:
- del weights[down_name]
- super().load_hf_weights(weights)
+ # [E, H, 2I] -> [E, 2I, H]
+ weights[gate_up_name] = weights[gate_up_name].transpose(1, 2).contiguous()
+ if down_name in weights:
+ # [E, I, H] -> [E, H, I]
+ weights[down_name] = weights[down_name].transpose(1, 2).contiguous()
diff --git a/lightllm/models/qwen3next/layer_infer/transformer_layer_infer.py b/lightllm/models/qwen3next/layer_infer/transformer_layer_infer.py
index d97cc3c12f..a1e8e63a6d 100644
--- a/lightllm/models/qwen3next/layer_infer/transformer_layer_infer.py
+++ b/lightllm/models/qwen3next/layer_infer/transformer_layer_infer.py
@@ -136,6 +136,7 @@ def _moe_ffn_tp(
use_grouped_topk=False,
topk_group=None,
num_expert_group=None,
+ infer_state=infer_state,
shared_expert_gate=shared_expert_gate,
)
hidden_states = hidden_states.view(num_tokens, hidden_dim)
@@ -157,6 +158,7 @@ def _moe_ffn_edp(
topk_group=None,
num_expert_group=None,
is_prefill=infer_state.is_prefill,
+ infer_state=infer_state,
)
ep_output = ep_output.view(token_num, hidden_dim)
ep_output.add_(shared_expert_out)
diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py
index e369d8257e..e85d083075 100644
--- a/lightllm/server/api_cli.py
+++ b/lightllm/server/api_cli.py
@@ -1,8 +1,7 @@
import argparse
-def make_argument_parser() -> argparse.ArgumentParser:
- parser = argparse.ArgumentParser()
+def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser:
parser.add_argument(
"--run_mode",
@@ -269,8 +268,6 @@ def make_argument_parser() -> argparse.ArgumentParser:
help="Whether or not to allow for custom models defined on the Hub in their own modeling files.",
)
parser.add_argument("--detail_log", action="store_true", help="enable to print input infos in requests.")
- parser.add_argument("--disable_log_stats", action="store_true", help="disable logging throughput stats.")
- parser.add_argument("--log_stats_interval", type=int, default=10, help="log stats interval in second.")
parser.add_argument(
"--disable_shm_warning",
action="store_true",
@@ -453,6 +450,11 @@ def make_argument_parser() -> argparse.ArgumentParser:
default=3686400, # 8294400 is 4k, 3686400 is 2k
help="maximum allowed pixel count for one image before resize preprocessing",
)
+ parser.add_argument(
+ "--disable_image_resize",
+ action="store_true",
+ help="disable automatic resize for images exceeding --max_image_pixels (enabled by default)",
+ )
parser.add_argument(
"--embed_cache_storage_size",
type=float,
@@ -466,7 +468,11 @@ def make_argument_parser() -> argparse.ArgumentParser:
default=None,
help="the data type of the model weight",
)
- parser.add_argument("--return_all_prompt_logprobs", action="store_true", help="return all prompt tokens logprobs")
+ parser.add_argument(
+ "--enable_prompt_logprobs",
+ action="store_true",
+ help="enable prompt top-k logprobs capture",
+ )
parser.add_argument("--use_reward_model", action="store_true", help="use reward model")
@@ -497,13 +503,6 @@ def make_argument_parser() -> argparse.ArgumentParser:
)
parser.add_argument("--visual_tp", type=int, default=1, help="number of tensort parallel instances for ViT")
parser.add_argument("--visual_dp", type=int, default=1, help="number of data parallel instances for ViT")
- parser.add_argument(
- "--visual_nccl_ports",
- nargs="+",
- type=int,
- default=None,
- help="List of NCCL ports to build a distributed environment for Vit, e.g., 29500 29501 29502",
- )
parser.add_argument(
"--visual_rpyc_port",
type=int,
@@ -523,13 +522,6 @@ def make_argument_parser() -> argparse.ArgumentParser:
help="Tensor parallel size for audio encoder (only 1 is supported; use audio_dp to scale)",
)
parser.add_argument("--audio_dp", type=int, default=1, help="Data parallel replicas for audio encoder")
- parser.add_argument(
- "--audio_nccl_ports",
- nargs="+",
- type=int,
- default=None,
- help="NCCL ports per audio DP group; if omitted, auto-allocated in api_start (reserved until audio_tp>1)",
- )
parser.add_argument(
"--audio_infer_batch_size",
type=int,
@@ -769,6 +761,19 @@ def make_argument_parser() -> argparse.ArgumentParser:
parser.add_argument(
"--disk_cache_storage_size", type=float, default=10, help="""The capacity of disk cache. GB used."""
)
+ parser.add_argument(
+ "--enable_rl",
+ action="store_true",
+ default=False,
+ help="""enable RL control plane (HTTP APIs, router rl_rpyc, model RlBackendOps).
+ When disabled (default), RL routes/services are not started.""",
+ )
+ parser.add_argument(
+ "--enable_torch_memory_saver",
+ action="store_true",
+ help="""enable torch memory saver, which is used for release_memory and resume_memory during RL training.""",
+ )
+ parser.add_argument("--enable_weight_cpu_backup", action="store_true", help="""enable weight cpu backup.""")
parser.add_argument(
"--disk_cache_dir",
type=str,
@@ -842,6 +847,12 @@ def make_argument_parser() -> argparse.ArgumentParser:
If the op is not implemented for the platform and the hardware support triton,
it will use triton implementation.""",
)
+ parser.add_argument(
+ "--enable_return_routed_experts",
+ action="store_true",
+ default=False,
+ help="Enable returning routed expert indices for MoE models (R3 feature).",
+ )
parser.add_argument(
"--enable_profiling",
type=str,
@@ -858,3 +869,7 @@ def make_argument_parser() -> argparse.ArgumentParser:
A NVTX range named 'LIGHTLLM_PROFILE' will be added within the profiling range.""",
)
return parser
+
+
+def make_argument_parser() -> argparse.ArgumentParser:
+ return add_cli_args(argparse.ArgumentParser())
diff --git a/lightllm/server/api_http.py b/lightllm/server/api_http.py
index e127f0931c..52144af123 100755
--- a/lightllm/server/api_http.py
+++ b/lightllm/server/api_http.py
@@ -34,7 +34,7 @@
import uuid
from PIL import Image
import multiprocessing as mp
-from typing import AsyncGenerator, Union
+from typing import Any, AsyncGenerator, Union
from typing import Callable
from lightllm.server import TokenLoad
from fastapi import BackgroundTasks, FastAPI, Request, WebSocket, WebSocketDisconnect
@@ -50,6 +50,7 @@
from lightllm.utils.error_utils import ClientDisconnected, ServerBusyError
from lightllm.server.metrics.manager import MetricClient
from lightllm.utils.envs_utils import get_unique_server_name
+from lightllm.utils.shm_port_args import get_shm_port_args
from dataclasses import dataclass
from .api_openai import chat_completions_impl, completions_impl
@@ -94,7 +95,7 @@ def set_args(self, args: StartArgs):
setproctitle.setproctitle(f"lightllm::{get_unique_server_name()}::api_server")
if args.run_mode == "pd_master":
- self.metric_client = MetricClient(args.metric_port)
+ self.metric_client = MetricClient(get_shm_port_args().metric_port)
self.httpserver_manager = HttpServerManagerForPDMaster(
args=args,
)
@@ -103,7 +104,7 @@ def set_args(self, args: StartArgs):
SamplingParams.load_generation_cfg(args.model_dir)
CompletionRequest.load_generation_cfg(args.model_dir)
ChatCompletionRequest.load_generation_cfg(args.model_dir)
- self.metric_client = MetricClient(args.metric_port)
+ self.metric_client = MetricClient(get_shm_port_args().metric_port)
self.httpserver_manager = HttpServerManager(args=args)
dp_size_in_node = max(1, args.dp // args.nnodes) # 兼容多机纯tp的运行模式,这时候 1 // 2 == 0, 需要兼容
self.shared_token_load = TokenLoad(f"{get_unique_server_name()}_shared_token_load", dp_size_in_node)
@@ -187,6 +188,22 @@ def get_model_name():
return {"model_name": g_objs.args.model_name}
+@app.get("/get_server_info")
+@app.post("/get_server_info")
+def get_server_info():
+ # 将 StartArgs 转换为字典格式
+ from dataclasses import asdict
+
+ server_info: dict[str, Any] = asdict(g_objs.args)
+ return {**server_info}
+
+
+@app.get("/get_weight_version")
+@app.post("/get_weight_version")
+def get_weight_version():
+ return {"weight_version": g_objs.args.weight_version}
+
+
@app.get("/healthz", summary="Check server health")
@app.get("/health", summary="Check server health")
@app.head("/health", summary="Check server health")
@@ -410,6 +427,12 @@ async def metrics() -> Response:
return response
+# RL 控制面接口(abort / pause / flush / memory / weight update),见 api_http_rl.py
+from .api_http_rl import router as rl_router
+
+app.include_router(rl_router)
+
+
@app.websocket("/pd_register")
async def register_and_keep_alive(websocket: WebSocket):
await websocket.accept()
diff --git a/lightllm/server/api_http_rl.py b/lightllm/server/api_http_rl.py
new file mode 100644
index 0000000000..f4ead5031b
--- /dev/null
+++ b/lightllm/server/api_http_rl.py
@@ -0,0 +1,143 @@
+"""RL control-plane HTTP APIs.
+
+需要启动参数 ``--enable_rl``。供 RL / 在线训推一体场景使用,不走普通 generate 路径。
+
+调用链:
+ HTTP → HttpServerManager → HttpRlController
+ → (多数) RlOpReq → Router → Model RlBackendOps
+
+路由在模块级 ``router`` 上注册,由 ``api_http`` ``include_router`` 挂载。
+``g_objs`` / ``create_error_response`` 在 handler 内懒导入,避免与 api_http 循环依赖。
+"""
+
+from http import HTTPStatus
+
+from fastapi import APIRouter, Request
+from fastapi.responses import JSONResponse, Response
+
+from lightllm.server.io_struct import (
+ AbortReq,
+ DestroyWeightsUpdateGroupReq,
+ FlushCacheReq,
+ InitWeightsUpdateGroupReq,
+ ReleaseMemoryReq,
+ ResumeMemoryReq,
+ RlOpRsp,
+ UpdateWeightsFromDistributedReq,
+ UpdateWeightsFromIPCReq,
+ UpdateWeightsFromTensorReq,
+)
+from lightllm.utils.log_utils import init_logger
+
+logger = init_logger(__name__)
+
+router = APIRouter()
+
+
+async def handle_request_common(request_obj, handler):
+ from .api_http import create_error_response
+
+ try:
+ ret: RlOpRsp = await handler(request_obj)
+ if ret.success:
+ return JSONResponse({"success": ret.success, "message": ret.msg}, status_code=200)
+ else:
+ return create_error_response(HTTPStatus.BAD_REQUEST, ret.msg)
+ except Exception as e:
+ logger.error("handle_request_common (%s) error occurred: %s", str(request_obj), str(e), exc_info=True)
+ return create_error_response(HTTPStatus.EXPECTATION_FAILED, f"error: {str(e)}")
+
+
+@router.post("/abort_request")
+async def abort_request(request: AbortReq, raw_request: Request):
+ """Abort a request."""
+ from .api_http import create_error_response, g_objs
+
+ try:
+ success, msg = await g_objs.httpserver_manager.abort_request(request)
+ if not success:
+ return create_error_response(HTTPStatus.REQUEST_TIMEOUT, msg, err_type="AbortRequestTimeout")
+ return Response(status_code=200)
+ except Exception as e:
+ logger.error("abort_request error occurred: %s", str(e), exc_info=True)
+ return create_error_response(HTTPStatus.EXPECTATION_FAILED, f"error: {str(e)}")
+
+
+@router.post("/init_weights_update_group")
+async def init_weights_update_group(request: InitWeightsUpdateGroupReq, raw_request: Request):
+ """Init weights update group."""
+ from .api_http import g_objs
+
+ return await handle_request_common(request, g_objs.httpserver_manager.init_weights_update_group)
+
+
+@router.post("/destroy_weights_update_group")
+async def destroy_weights_update_group(request: DestroyWeightsUpdateGroupReq, raw_request: Request):
+ """Destroy weights update group."""
+ from .api_http import g_objs
+
+ return await handle_request_common(request, g_objs.httpserver_manager.destroy_weights_update_group)
+
+
+@router.post("/update_weights_from_distributed")
+async def update_weights_from_distributed(request: UpdateWeightsFromDistributedReq, raw_request: Request):
+ """Update model parameter from distributed online."""
+ from .api_http import g_objs
+
+ return await handle_request_common(request, g_objs.httpserver_manager.update_weights_from_distributed)
+
+
+@router.post("/update_weights_from_tensor")
+async def update_weights_from_tensor(request: UpdateWeightsFromTensorReq, raw_request: Request):
+ """Update model parameter from distributed online."""
+ from .api_http import g_objs
+
+ return await handle_request_common(request, g_objs.httpserver_manager.update_weights_from_tensor)
+
+
+@router.post("/update_weights_from_ipc")
+async def update_weights_from_ipc(request: UpdateWeightsFromIPCReq, raw_request: Request):
+ from .api_http import g_objs
+
+ return await handle_request_common(request, g_objs.httpserver_manager.update_weights_from_ipc)
+
+
+@router.post("/flush_cache")
+@router.get("/flush_cache")
+async def flush_cache():
+ """Flush the radix cache."""
+ from .api_http import g_objs
+
+ return await handle_request_common(FlushCacheReq(), g_objs.httpserver_manager.flush_cache)
+
+
+@router.post("/pause_generation")
+async def pause_generation():
+ from .api_http import g_objs
+
+ await g_objs.httpserver_manager.pause_generation()
+ return Response(content="Generation paused successfully.", status_code=200)
+
+
+@router.post("/continue_generation")
+async def continue_generation():
+ from .api_http import g_objs
+
+ await g_objs.httpserver_manager.continue_generation()
+ return Response(content="Generation continued successfully.", status_code=200)
+
+
+@router.get("/release_memory_occupation")
+@router.post("/release_memory_occupation")
+async def release_memory_occupation(request: ReleaseMemoryReq):
+ from .api_http import g_objs
+
+ return await handle_request_common(request, g_objs.httpserver_manager.release_memory_occupation)
+
+
+@router.get("/resume_memory_occupation")
+@router.post("/resume_memory_occupation")
+async def resume_memory_occupation(request: ResumeMemoryReq):
+ from .api_http import g_objs
+
+ return await handle_request_common(request, g_objs.httpserver_manager.resume_memory_occupation)
diff --git a/lightllm/server/api_lightllm.py b/lightllm/server/api_lightllm.py
index 39a5808aab..6a0abe81be 100644
--- a/lightllm/server/api_lightllm.py
+++ b/lightllm/server/api_lightllm.py
@@ -35,6 +35,9 @@ async def lightllm_generate(request: Request, httpserver_manager: HttpServerMana
prompt = request_dict.pop("inputs")
sample_params_dict = request_dict["parameters"]
return_details = sample_params_dict.pop("return_details", False)
+ return_routed_experts = sample_params_dict.pop(
+ "return_routed_experts", httpserver_manager.args.enable_return_routed_experts
+ )
sampling_params = SamplingParams()
sampling_params.init(tokenizer=httpserver_manager.tokenizer, **sample_params_dict)
sampling_params.verify()
@@ -47,43 +50,45 @@ async def lightllm_generate(request: Request, httpserver_manager: HttpServerMana
final_output_dict = collections.defaultdict(list)
count_output_tokens_dict = collections.defaultdict(lambda: 0)
tokens_dict = collections.defaultdict(list)
+ logprobs_dict = collections.defaultdict(list)
finish_reason_dict = {}
prompt_logprobs = None
prompt_tokens = 0
prompt_token_ids = None
is_first_metadata = True
input_usage = None
+ routed_experts_data = None
async for sub_req_id, request_output, metadata, finish_status in results_generator:
- # when set "--return_all_prompt_logprobs", the first token metadata will contains
- # prompt_logprobs and prompt_token_ids
if is_first_metadata:
- prompt_logprobs = metadata.get("prompt_logprobs", None)
- prompt_token_ids = metadata.get("prompt_token_ids", None)
prompt_tokens = metadata.get("prompt_tokens", 0)
input_usage = metadata.get("input_usage", None)
- if prompt_logprobs is not None:
- del metadata["prompt_logprobs"]
- if prompt_token_ids is not None:
- del metadata["prompt_token_ids"]
if input_usage is not None:
del metadata["input_usage"]
is_first_metadata = False
+ if "prompt_logprobs" in metadata:
+ prompt_logprobs = metadata.pop("prompt_logprobs")
+ prompt_token_ids = metadata.pop("prompt_token_ids", None)
+
count_output_tokens_dict[sub_req_id] += 1
final_output_dict[sub_req_id].append(request_output)
+ logprobs_dict[sub_req_id].append(metadata.pop("logprobs"))
if return_details:
metadata["text"] = request_output
tokens_dict[sub_req_id].append(metadata)
if finish_status.is_finished():
finish_reason_dict[sub_req_id] = finish_status
+ if "routed_experts" in metadata:
+ routed_experts_data = metadata["routed_experts"]
n = sampling_params.n
sub_ids = list(final_output_dict.keys())[:n]
final_output_list = ["".join(final_output_dict[sub_id]) for sub_id in sub_ids]
count_output_tokens_list = [count_output_tokens_dict[sub_id] for sub_id in sub_ids]
finish_reson_list = [finish_reason_dict[sub_id].get_finish_reason() for sub_id in sub_ids]
tokens_list = [tokens_dict[sub_id] for sub_id in sub_ids]
+ logprobs_list = [logprobs_dict[sub_id] for sub_id in sub_ids]
only_one = len(sub_ids) == 1
ret_data_format = lambda data_list: data_list[0] if only_one else data_list
@@ -96,12 +101,15 @@ async def lightllm_generate(request: Request, httpserver_manager: HttpServerMana
}
if return_details:
ret["tokens"] = ret_data_format(tokens_list)
+ ret["logprobs"] = ret_data_format(logprobs_list)
if prompt_token_ids is not None:
ret["prompt_token_ids"] = prompt_token_ids
if prompt_logprobs is not None:
ret["prompt_logprobs"] = prompt_logprobs
if input_usage is not None:
ret["input_usage"] = input_usage
+ if return_routed_experts and routed_experts_data is not None:
+ ret["routed_experts"] = routed_experts_data
return Response(content=json.dumps(ret, ensure_ascii=False).encode("utf-8"))
@@ -112,6 +120,7 @@ async def lightllm_generate_stream(request: Request, httpserver_manager: HttpSer
prompt = request_dict.pop("inputs")
sample_params_dict = request_dict["parameters"]
_ = sample_params_dict.pop("return_details", False)
+ _ = sample_params_dict.pop("return_routed_experts", None)
sampling_params = SamplingParams()
sampling_params.init(tokenizer=httpserver_manager.tokenizer, **sample_params_dict)
sampling_params.verify()
@@ -145,6 +154,10 @@ async def stream_results() -> AsyncGenerator[bytes, None]:
"details": None,
"input_usage": input_usage,
}
+ ret["token"]["logprobs"] = metadata["logprobs"]
+ if "prompt_logprobs" in metadata:
+ ret["prompt_logprobs"] = metadata["prompt_logprobs"]
+ ret["prompt_token_ids"] = metadata.get("prompt_token_ids")
yield ("data:" + json.dumps(ret, ensure_ascii=False) + "\n\n").encode("utf-8")
diff --git a/lightllm/server/api_server.py b/lightllm/server/api_server.py
index 6e04d5d47e..5306ecb698 100755
--- a/lightllm/server/api_server.py
+++ b/lightllm/server/api_server.py
@@ -1,11 +1,22 @@
import torch
-from .api_cli import make_argument_parser
+from .api_cli import add_cli_args
+from lightllm.server.core.objs.start_args_type import StartArgs
+from lightllm.utils.log_utils import init_logger
+
+logger = init_logger(__name__)
-if __name__ == "__main__":
- torch.multiprocessing.set_start_method("spawn") # this code will not be ok for settings to fork to subprocess
- parser = make_argument_parser()
- args = parser.parse_args()
- from .api_start import pd_master_start, normal_or_p_d_start, visual_only_start, config_server_start
+
+def launch_server(args: StartArgs):
+ from .api_start import pd_master_start, normal_or_p_d_start, config_server_start, visual_only_start
+
+ try:
+ # this code will not be ok for settings to fork to subprocess
+ torch.multiprocessing.set_start_method("spawn")
+ except RuntimeError as e:
+ logger.warning(f"Failed to set start method: {e}")
+ except Exception as e:
+ logger.error(f"Failed to set start method: {e}")
+ raise e
if args.run_mode == "pd_master":
pd_master_start(args)
@@ -15,3 +26,13 @@
visual_only_start(args)
else:
normal_or_p_d_start(args)
+
+
+if __name__ == "__main__":
+ from argparse import ArgumentParser
+
+ parser = ArgumentParser()
+ add_cli_args(parser)
+ args = parser.parse_args()
+
+ launch_server(StartArgs(**vars(args)))
diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py
index bec1be05b8..134c6623c4 100644
--- a/lightllm/server/api_start.py
+++ b/lightllm/server/api_start.py
@@ -1,3 +1,4 @@
+import multiprocessing as mp
import os
import sys
import time
@@ -5,19 +6,21 @@
import subprocess
import signal
import math
-from lightllm.utils.net_utils import alloc_can_use_network_port, PortLocker
from lightllm.utils.start_utils import process_manager, kill_recursive
from .metrics.manager import start_metric_manager
from .embed_cache.manager import start_cache_manager
from lightllm.utils.log_utils import init_logger
from lightllm.utils.envs_utils import set_env_start_args, set_unique_server_name, get_unique_server_name
from lightllm.utils.envs_utils import get_lightllm_gunicorn_keep_alive
+from lightllm.utils.shm_port_args import get_shm_port_args
+from lightllm.utils.net_utils import validate_ports
from .detokenization.manager import start_detokenization_process
from .router.manager import start_router_process
from lightllm.utils.process_check import is_process_active
from lightllm.utils.multinode_utils import send_and_receive_node_ip
from lightllm.utils.redis_utils import start_redis_service
from lightllm.utils.shm_size_check import check_recommended_shm_size
+from lightllm.server.core.objs.start_args_type import StartArgs
from lightllm.utils.config_utils import (
has_audio_module,
has_vision_module,
@@ -61,9 +64,31 @@ def signal_handler(sig, frame):
process_manager.terminate_all_processes()
logger.info("All processes have been terminated gracefully.")
sys.exit(0)
+ elif sig == signal.SIGHUP:
+ logger.info("Received SIGHUP (terminal closed), shutting down gracefully...")
+ if http_server_process and http_server_process.poll() is None:
+ http_server_process.send_signal(signal.SIGTERM)
+
+ start_time = time.time()
+ while (time.time() - start_time) < 60:
+ if not is_process_active(http_server_process.pid):
+ logger.info("httpserver exit")
+ break
+ time.sleep(1)
+
+ if time.time() - start_time < 60:
+ logger.info("HTTP server has exited gracefully")
+ else:
+ logger.warning("HTTP server did not exit in time, killing it...")
+ kill_recursive(http_server_process)
+
+ process_manager.terminate_all_processes()
+ logger.info("All processes have been terminated gracefully due to terminal closure.")
+ sys.exit(0)
signal.signal(signal.SIGTERM, signal_handler)
signal.signal(signal.SIGINT, signal_handler)
+ signal.signal(signal.SIGHUP, signal_handler)
logger.info(f"start process pid {os.getpid()}")
if http_server_process:
@@ -71,10 +96,12 @@ def signal_handler(sig, frame):
return
-def normal_or_p_d_start(args):
- from lightllm.server.core.objs.start_args_type import StartArgs
+def _set_envs_and_config(args: StartArgs):
+ mp.set_start_method("spawn", force=True)
- args: StartArgs = args
+
+def _launch_subprocesses(args: StartArgs):
+ _set_envs_and_config(args)
auto_set_max_req_total_len(args)
auto_set_fused_shared_experts(args)
@@ -145,12 +172,6 @@ def normal_or_p_d_start(args):
check_recommended_shm_size(args)
assert args.zmq_mode in ["tcp://", "ipc:///tmp/"]
- # 确保单机上多实列不冲突
- if args.zmq_mode == "ipc:///tmp/":
- zmq_mode = f"{args.zmq_mode}_{get_unique_server_name()}_"
- args.zmq_mode = None # args 的参数不能直接设置,只能先设置None,再设置才能成功
- args.zmq_mode = zmq_mode
- logger.info(f"zmq mode head: {args.zmq_mode}")
logger.info(f"use tgi api: {args.use_tgi_api}")
@@ -178,10 +199,6 @@ def normal_or_p_d_start(args):
if args.use_reward_model:
assert args.disable_dynamic_prompt_cache is True, "need add --disable_dynamic_prompt_cache"
assert args.disable_chunked_prefill is True, "need add --disable_chunked_prefill"
- if args.return_all_prompt_logprobs:
- assert args.disable_dynamic_prompt_cache is True, "need add --disable_dynamic_prompt_cache"
- assert args.disable_chunked_prefill is True, "need add --disable_chunked_prefill"
-
# FP8 KV cache mode checks
if args.llm_kv_type in ["fp8kv_sph", "fp8kv_spt"]:
assert (
@@ -212,12 +229,16 @@ def normal_or_p_d_start(args):
# mtp params check
if args.mtp_mode is not None:
- assert args.mtp_draft_model_dir is not None
+ if args.mtp_draft_model_dir is None:
+ args.mtp_draft_model_dir = [args.model_dir] * args.mtp_step
assert args.mtp_step > 0
else:
assert args.mtp_draft_model_dir is None
assert args.mtp_step == 0
+ # automatically set visual_dp based on visual_tp and tp
+ if args.visual_tp < args.tp and args.tp % args.visual_tp == 0:
+ args.visual_dp = args.tp // args.visual_tp
if args.afs_image_embed_dir is not None:
os.makedirs(args.afs_image_embed_dir, mode=0o777, exist_ok=True)
os.chmod(args.afs_image_embed_dir, 0o777)
@@ -339,72 +360,21 @@ def normal_or_p_d_start(args):
args.data_type = get_dtype(args.model_dir)
assert args.data_type in ["fp16", "float16", "bf16", "bfloat16", "fp32", "float32"]
- already_uesd_ports = [args.port]
- if args.nccl_port is not None:
- already_uesd_ports.append(args.nccl_port)
- if args.visual_nccl_ports is not None:
- already_uesd_ports.extend(args.visual_nccl_ports[: args.visual_dp])
- if not args.disable_audio and args.audio_nccl_ports is not None:
- already_uesd_ports.extend(args.audio_nccl_ports[: args.audio_dp])
-
- # 提前锁定端口,防止在单个机器上启动多个实列的时候,要到模型启动的时候才能
- # 捕获到端口设置冲突的问题
- ports_locker = PortLocker(already_uesd_ports)
- ports_locker.lock_port()
-
- node_world_size = args.tp // args.nnodes
- can_use_ports = alloc_can_use_network_port(
- num=10 + node_world_size + args.visual_dp * args.visual_tp + args.visual_dp + args.audio_dp,
- used_ports=already_uesd_ports,
- )
- logger.info(f"alloced ports: {can_use_ports}")
- (
- nccl_port,
- router_port,
- router_profiler_port,
- detokenization_port,
- http_server_port,
- visual_port,
- audio_port,
- cache_port,
- metric_port,
- multi_level_kv_cache_port,
- ) = can_use_ports[0:10]
- can_use_ports = can_use_ports[10:]
-
- if args.visual_nccl_ports is None:
- args.visual_nccl_ports = can_use_ports[: args.visual_dp]
- can_use_ports = can_use_ports[args.visual_dp :]
- else:
- args.visual_nccl_ports = args.visual_nccl_ports[: args.visual_dp]
+ set_unique_server_name(args)
+
+ # 确保单机上多实列不冲突
+ if args.zmq_mode == "ipc:///tmp/":
+ zmq_mode = f"{args.zmq_mode}_{get_unique_server_name()}_"
+ args.zmq_mode = None # args 的参数不能直接设置,只能先设置None,再设置才能成功
+ args.zmq_mode = zmq_mode
+ logger.info(f"zmq mode head: {args.zmq_mode}")
- if args.audio_nccl_ports is None:
- args.audio_nccl_ports = can_use_ports[: args.audio_dp]
- can_use_ports = can_use_ports[args.audio_dp :]
- else:
- args.audio_nccl_ports = args.audio_nccl_ports[: args.audio_dp]
-
- # 将申请好的端口放入args参数中
- if args.nccl_port is None:
- args.nccl_port = nccl_port
- args.router_port = router_port
- args.router_profiler_port = router_profiler_port
- args.detokenization_port = detokenization_port
- args.http_server_port = http_server_port
- args.visual_port = visual_port
- args.audio_port = audio_port
- args.cache_port = cache_port
- args.metric_port = metric_port
- args.multi_level_kv_cache_port = multi_level_kv_cache_port
- # 申请在 p d 分离模式下,会用的端口
- args.pd_node_infer_rpyc_ports = can_use_ports[0:node_world_size]
# p d 分离模式下用于标识节点的id
args.pd_node_id = uuid.uuid4().int
# p d 分离模式下,decode节点的调度间隙是0
if args.run_mode == "decode":
args.router_max_wait_tokens = 0
- send_and_receive_node_ip(args) # 多机用于收发node ip
# dp 必须 > 1
if args.enable_dp_prompt_cache_fetch and args.dp <= 1:
args.enable_dp_prompt_cache_fetch = False
@@ -415,11 +385,19 @@ def normal_or_p_d_start(args):
auto_configure_allreduce_flags_from_args(args)
+ # 校验用户已设置端口冲突(对齐原 PortManager 启动检查范围)
+ ports_to_check = [args.port, args.multinode_httpmanager_port, args.multinode_router_gloo_port]
+ if args.node_rank == 0 and args.nccl_port is not None:
+ ports_to_check.append(args.nccl_port)
+ validate_ports(ports_to_check)
+
+ set_env_start_args(args)
+ get_shm_port_args(create=True)
+ # 多机用于收发node ip, 这个地方修改了args env,所以需要重新设置一下。
+ send_and_receive_node_ip(args)
set_env_start_args(args)
logger.info(f"all start args:{args}")
- ports_locker.release_port()
-
if args.enable_multimodal:
process_manager.start_submodule_processes(
start_funcs=[
@@ -429,7 +407,6 @@ def normal_or_p_d_start(args):
)
if not args.disable_vision:
-
if not args.visual_use_proxy_mode:
from .visualserver.manager import start_visual_process
@@ -490,13 +467,19 @@ def normal_or_p_d_start(args):
],
)
+ return process_manager
+
+
+def normal_or_p_d_start(args: StartArgs):
+ process_manager = _launch_subprocesses(args)
+
# 启动 Hypercorn
command = [
"hypercorn",
"--workers",
f"{args.httpserver_workers}",
"--bind",
- f"{args.host}:{args.port}",
+ f"{args.host}:{get_shm_port_args().port}",
"--log-level",
"info",
"--access-logfile",
@@ -525,7 +508,8 @@ def normal_or_p_d_start(args):
return
-def pd_master_start(args):
+def pd_master_start(args: StartArgs):
+ _set_envs_and_config(args)
set_unique_server_name(args)
if args.run_mode != "pd_master":
return
@@ -541,19 +525,11 @@ def pd_master_start(args):
args.pd_node_id = 0
logger.info(f"use tgi api: {args.use_tgi_api}")
- logger.info(f"all start args:{args}")
-
- can_use_ports = alloc_can_use_network_port(
- num=1,
- used_ports=[
- args.port,
- ],
- )
- metric_port = can_use_ports[0]
-
- args.metric_port = metric_port
+ validate_ports([args.port])
set_env_start_args(args)
+ get_shm_port_args(create=True)
+ logger.info(f"all start args:{args}")
process_manager.start_submodule_processes(
start_funcs=[
@@ -567,7 +543,7 @@ def pd_master_start(args):
"--workers",
"1",
"--bind",
- f"{args.host}:{args.port}",
+ f"{args.host}:{get_shm_port_args().port}",
"--log-level",
"info",
"--access-logfile",
@@ -594,16 +570,12 @@ def visual_only_start(args):
from lightllm.server.core.objs.start_args_type import StartArgs
args: StartArgs = args
+ _set_envs_and_config(args)
if args.afs_image_embed_dir is not None:
os.makedirs(args.afs_image_embed_dir, mode=0o777, exist_ok=True)
os.chmod(args.afs_image_embed_dir, 0o777)
- already_uesd_ports = []
- already_uesd_ports.append(args.visual_rpyc_port)
- can_use_ports = alloc_can_use_network_port(
- num=5 + args.visual_dp * args.visual_tp + args.visual_dp,
- used_ports=already_uesd_ports,
- )
+ set_unique_server_name(args)
if args.visual_gpu_ids is None:
args.visual_gpu_ids = list(range(args.visual_dp * args.visual_tp))
@@ -615,15 +587,15 @@ def visual_only_start(args):
args.data_type = get_dtype(args.model_dir)
assert args.data_type in ["fp16", "float16", "bf16", "bfloat16", "fp32", "float32"]
- logger.info(f"alloced ports: {can_use_ports}")
-
- args.visual_nccl_ports = can_use_ports[: args.visual_dp]
- can_use_ports = can_use_ports[args.visual_dp :]
args.visual_node_id = uuid.uuid4().int
- logger.info(f"all start args:{args}")
-
+ ports_to_check = []
+ if args.visual_rpyc_port is not None:
+ ports_to_check.append(args.visual_rpyc_port)
+ validate_ports(ports_to_check)
set_env_start_args(args)
+ get_shm_port_args(create=True)
+ logger.info(f"all start args:{args}")
from .visualserver.visual_only_manager import start_visual_process
@@ -651,19 +623,23 @@ def config_server_start(args):
if args.run_mode != "config_server":
return
+ ports_to_check = [args.config_server_port]
+ if args.config_server_visual_redis_port is not None:
+ ports_to_check.append(args.config_server_visual_redis_port)
+ validate_ports(ports_to_check)
+ set_env_start_args(args)
+ get_shm_port_args(create=True)
logger.info(f"all start args:{args}")
if args.config_server_visual_redis_port is not None:
start_redis_service(args)
- set_env_start_args(args)
-
command = [
"hypercorn",
"--workers",
"1",
"--bind",
- f"{args.config_server_host}:{args.config_server_port}",
+ f"{args.config_server_host}:{get_shm_port_args().config_server_port}",
"--log-level",
"info",
"--access-logfile",
diff --git a/lightllm/server/api_tgi.py b/lightllm/server/api_tgi.py
index f4a7cf6a5a..2c69623238 100755
--- a/lightllm/server/api_tgi.py
+++ b/lightllm/server/api_tgi.py
@@ -76,26 +76,20 @@ async def tgi_generate_impl(request: Request, httpserver_manager: HttpServerMana
final_output_dict = collections.defaultdict(list)
count_output_tokens_dict = collections.defaultdict(lambda: 0)
tokens_dict = collections.defaultdict(list)
+ logprobs_dict = collections.defaultdict(list)
finish_status_dict = {}
prompt_logprobs = None
prompt_token_ids = None
- is_first_metadata = True
best_score = -float("inf")
best_sub_id = 0
async for sub_req_id, request_output, metadata, finish_status in results_generator:
- # when set "--return_all_prompt_logprobs", the first token metadata will contains
- # prompt_logprobs and prompt_token_ids
- if is_first_metadata:
- prompt_logprobs = metadata.get("prompt_logprobs", None)
- prompt_token_ids = metadata.get("prompt_token_ids", None)
- if prompt_logprobs is not None:
- del metadata["prompt_logprobs"]
- if prompt_token_ids is not None:
- del metadata["prompt_token_ids"]
- is_first_metadata = False
+ if "prompt_logprobs" in metadata:
+ prompt_logprobs = metadata.pop("prompt_logprobs")
+ prompt_token_ids = metadata.pop("prompt_token_ids", None)
count_output_tokens_dict[sub_req_id] += 1
final_output_dict[sub_req_id].append(request_output)
+ logprobs_dict[sub_req_id].append(metadata.pop("logprobs"))
if return_details:
metadata["text"] = request_output
tokens_dict[sub_req_id].append(metadata)
@@ -132,6 +126,7 @@ async def tgi_generate_impl(request: Request, httpserver_manager: HttpServerMana
ret["prompt_token_ids"] = prompt_token_ids
if prompt_logprobs is not None:
ret["prompt_logprobs"] = prompt_logprobs
+ ret["logprobs"] = logprobs_dict[sub_id]
assert ret is not None
if return_details:
ret["details"]["beam_sequences"] = beam_sequences
@@ -177,6 +172,7 @@ async def stream_results() -> AsyncGenerator[bytes, None]:
"finish_reason": finish_status.get_finish_reason(),
"details": None,
}
+ ret["token"]["logprobs"] = metadata["logprobs"]
final_output.append(request_output)
if ret["finished"]:
ret["generated_text"] = "".join(final_output)
@@ -186,6 +182,9 @@ async def stream_results() -> AsyncGenerator[bytes, None]:
"finish_reason": finish_status.get_finish_reason(),
"prompt_tokens": metadata.get("prompt_tokens", 0),
}
+ if "prompt_logprobs" in metadata:
+ ret["prompt_logprobs"] = metadata["prompt_logprobs"]
+ ret["prompt_token_ids"] = metadata.get("prompt_token_ids")
yield "data:" + json.dumps(ret, ensure_ascii=False) + "\n\n"
diff --git a/lightllm/server/audioserver/manager.py b/lightllm/server/audioserver/manager.py
index efe24c53e3..027c5b7aa8 100644
--- a/lightllm/server/audioserver/manager.py
+++ b/lightllm/server/audioserver/manager.py
@@ -19,6 +19,7 @@
from lightllm.utils.graceful_utils import graceful_registry
from lightllm.utils.process_check import start_parent_check_thread
from lightllm.utils.envs_utils import get_unique_server_name
+from lightllm.utils.shm_port_args import get_shm_port_args
from rpyc.utils.classic import obtain
@@ -31,18 +32,19 @@ def __init__(
args: StartArgs,
):
self.args = args
+ ports = get_shm_port_args()
context = zmq.Context(2)
if args.enable_cpu_cache:
self.send_to_next_module = context.socket(zmq.PUSH)
- self.send_to_next_module.connect(f"{args.zmq_mode}127.0.0.1:{args.multi_level_kv_cache_port}")
+ self.send_to_next_module.connect(f"{args.zmq_mode}127.0.0.1:{ports.multi_level_kv_cache_port}")
else:
self.send_to_next_module = context.socket(zmq.PUSH)
- self.send_to_next_module.connect(f"{args.zmq_mode}127.0.0.1:{args.router_port}")
+ self.send_to_next_module.connect(f"{args.zmq_mode}127.0.0.1:{ports.router_port}")
self.zmq_recv_socket = context.socket(zmq.PULL)
- self.zmq_recv_socket.bind(f"{args.zmq_mode}127.0.0.1:{args.audio_port}")
- self.cache_client = rpyc.connect("localhost", args.cache_port, config={"allow_pickle": True})
+ self.zmq_recv_socket.bind(f"{args.zmq_mode}127.0.0.1:{ports.audio_port}")
+ self.cache_client = rpyc.connect("localhost", ports.cache_port, config={"allow_pickle": True})
self.cache_client._channel.stream.sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1)
self.model_weightdir = args.model_dir
self.audio_dp = args.audio_dp
@@ -59,6 +61,8 @@ async def wait_to_model_ready(self):
self.model_rpcs[dp_rank_id].append(rpc_model)
init_model_ret = []
+ ports = get_shm_port_args()
+ audio_nccl_ports = ports.audio_nccl_ports
for dp_rank_id in range(self.audio_dp):
for tp_rank_id in range(self.audio_tp):
device_id = self.args.audio_gpu_ids[dp_rank_id * self.audio_tp + tp_rank_id]
@@ -66,11 +70,11 @@ async def wait_to_model_ready(self):
"weight_dir": self.model_weightdir,
"device_id": device_id,
"audio_tp": self.audio_tp,
- "cache_port": self.args.cache_port,
+ "cache_port": ports.cache_port,
"tp_rank_id": tp_rank_id,
"dp_rank_id": dp_rank_id,
"data_type": self.args.data_type,
- "audio_nccl_port": self.args.audio_nccl_ports[dp_rank_id],
+ "audio_nccl_port": audio_nccl_ports[dp_rank_id],
"max_batch_size": max(self.infer_batch_size // self.audio_dp, 1),
}
init_model_ret.append(self.model_rpcs[dp_rank_id][tp_rank_id].init_model(kvargs))
diff --git a/lightllm/server/core/objs/logprob_utils.py b/lightllm/server/core/objs/logprob_utils.py
new file mode 100644
index 0000000000..fec92b82c8
--- /dev/null
+++ b/lightllm/server/core/objs/logprob_utils.py
@@ -0,0 +1,9 @@
+def logprob_info(tokenizer, token_id: int, logprob: float, rank: int):
+ decoded_token = None
+ if tokenizer is not None:
+ decoded_token = tokenizer.decode([token_id], skip_special_tokens=False)
+ return {
+ "logprob": float(logprob),
+ "rank": rank,
+ "decoded_token": decoded_token,
+ }
diff --git a/lightllm/server/core/objs/py_sampling_params.py b/lightllm/server/core/objs/py_sampling_params.py
index cbc63c898d..2514d9dacb 100644
--- a/lightllm/server/core/objs/py_sampling_params.py
+++ b/lightllm/server/core/objs/py_sampling_params.py
@@ -111,13 +111,18 @@ def __init__(
def load_generation_cfg(cls, weight_dir):
try:
generation_cfg = GenerationConfig.from_pretrained(weight_dir, trust_remote_code=True).to_dict()
- cls._do_sample = generation_cfg.get("do_sample", False)
- cls._presence_penalty = generation_cfg.get("presence_penalty", 0.0)
- cls._frequency_penalty = generation_cfg.get("frequency_penalty", 0.0)
- cls._repetition_penalty = generation_cfg.get("repetition_penalty", 1.0)
- cls._temperature = generation_cfg.get("temperature", 1.0)
- cls._top_p = generation_cfg.get("top_p", 1.0)
- cls._top_k = generation_cfg.get("top_k", -1)
+
+ def _cfg(key, default):
+ v = generation_cfg.get(key)
+ return v if v is not None else default
+
+ cls._do_sample = _cfg("do_sample", False)
+ cls._presence_penalty = _cfg("presence_penalty", 0.0)
+ cls._frequency_penalty = _cfg("frequency_penalty", 0.0)
+ cls._repetition_penalty = _cfg("repetition_penalty", 1.0)
+ cls._temperature = _cfg("temperature", 1.0)
+ cls._top_p = _cfg("top_p", 1.0)
+ cls._top_k = _cfg("top_k", -1)
cls._stop_sequences = generation_cfg.get("stop", None)
except:
pass
diff --git a/lightllm/server/core/objs/req.py b/lightllm/server/core/objs/req.py
index 7f2b697091..b812cacb6d 100644
--- a/lightllm/server/core/objs/req.py
+++ b/lightllm/server/core/objs/req.py
@@ -1,6 +1,7 @@
import os
import math
import ctypes
+import asyncio
import numpy as np
import time
from .sampling_params import SamplingParams
@@ -12,32 +13,48 @@
from lightllm.utils.envs_utils import get_env_start_args
from lightllm.utils.config_utils import is_linear_att_mixed_model
from lightllm.utils.kv_cache_utils import compute_token_list_hash
-from typing import List, Any, Union
+from typing import Any, Dict, List, Union
from lightllm.utils.log_utils import init_logger
+from .logprob_utils import logprob_info
+from .token_metadata import ReqFinalTokenMetadata
logger = init_logger(__name__)
class FinishStatus(ctypes.Structure):
+ """请求结束状态。API 侧通过 ``get_finish_reason()`` 映射为字符串。
+
+ - ``NO_FINISH``: 未结束
+ - ``FINISHED_STOP``: 正常停止(EOS / stop 序列等),finish_reason=``stop``
+ - ``FINISHED_LENGTH``: 达到 max_new_tokens 等长度上限,finish_reason=``length``
+ - ``FINISHED_ABORTED``: 客户端/调度主动 abort,finish_reason=``abort``
+ - ``FINISHED_ERROR``: 服务端内部错误导致无法继续生成,finish_reason=``error``。
+ 典型场景:PD 分离 decode 节点 KV 传输失败。与 abort 区分:非用户取消,
+ 而是传输/系统故障;若同一请求已 abort,应优先标 ``FINISHED_ABORTED``。
+ """
+
_pack_ = 4
_fields_ = [("status", ctypes.c_int)]
NO_FINISH = 0
FINISHED_STOP = 1
FINISHED_LENGTH = 2
+ FINISHED_ABORTED = 3
+ # 内部错误结束(如 PD KV 传输失败);见类文档。
+ FINISHED_ERROR = 4
def __init__(self, init_state=NO_FINISH):
self.status = init_state
def set_status(self, new_status):
- assert 0 <= new_status <= 2
+ assert 0 <= new_status <= 4
self.status = new_status
def get_status(self):
return self.status
def is_finished(self):
- return self.FINISHED_STOP <= self.status <= self.FINISHED_LENGTH
+ return self.FINISHED_STOP <= self.status <= self.FINISHED_ERROR
def is_stopped(self):
return self.status == self.FINISHED_STOP
@@ -45,11 +62,18 @@ def is_stopped(self):
def is_finished_length(self):
return self.status == self.FINISHED_LENGTH
+ def is_finished_error(self):
+ return self.status == self.FINISHED_ERROR
+
def get_finish_reason(self):
if self.status == self.FINISHED_STOP:
return "stop"
elif self.status == self.FINISHED_LENGTH:
return "length"
+ elif self.status == self.FINISHED_ABORTED:
+ return "abort"
+ elif self.status == self.FINISHED_ERROR:
+ return "error"
return None
@@ -266,17 +290,69 @@ def link_prompt_ids_shm_array(self):
def create_logprobs_shm_array(self):
service_uni_name = get_unique_server_name()
name = f"{service_uni_name}_shm_logprobs_{self.index_in_shm_mem}"
- self.shm_logprobs = ShmArray(name, (self.alloc_shm_numpy_len,), dtype=np.float32)
+ self.shm_logprobs = ShmArray(
+ name,
+ (self.alloc_shm_numpy_len,),
+ dtype=[("logprob", np.float32), ("rank", np.int32)],
+ )
self.shm_logprobs.create_shm()
+ # rank=-1 表示该位置没有请求或没有计算 rank 元信息。
+ self.shm_logprobs.arr["logprob"][:] = 0.0
+ self.shm_logprobs.arr["rank"][:] = -1
return
def link_logprobs_shm_array(self):
service_uni_name = get_unique_server_name()
name = f"{service_uni_name}_shm_logprobs_{self.index_in_shm_mem}"
- self.shm_logprobs = ShmArray(name, (self.alloc_shm_numpy_len,), dtype=np.float32)
+ self.shm_logprobs = ShmArray(
+ name,
+ (self.alloc_shm_numpy_len,),
+ dtype=[("logprob", np.float32), ("rank", np.int32)],
+ )
self.shm_logprobs.link_shm()
return
+ async def merge_final_token_metadata(
+ self,
+ metadata: Dict[str, Any],
+ tokenizer: Any,
+ enable_return_routed_experts: bool = False,
+ timeout: float = 60.0,
+ ) -> None:
+ """等待并读取 final token metadata,按需合并进 HTTP 输出 ``metadata``。
+
+ 仅在需要 ``prompt_logprobs`` / ``routed_experts`` 时执行;失败或超时
+ 不改动 ``metadata``(仅打 warning)。
+ """
+ # 阶段 1:判断本请求是否需要 final token metadata。
+ need_prompt_logprobs = self.sample_params.prompt_logprobs >= 0
+ if not (need_prompt_logprobs or enable_return_routed_experts):
+ return
+
+ # 阶段 2:等待 Infer 写完 metadata 并释放。
+ # Infer 在 dump 之后才会置 shm_infer_released=True,以此作为可读信号。
+ start_time = time.time()
+ while not self.shm_infer_released:
+ if time.time() - start_time > timeout:
+ logger.warning(f"wait final_token_metadata ready timeout, req_id={self.request_id}, timeout={timeout}s")
+ return
+ await asyncio.sleep(0.005)
+
+ # 阶段 3:从 shm 读取并解码(read 内部已尽量吞掉 shm 缺失等错误)。
+ try:
+ meta = ReqFinalTokenMetadata(self).read(tokenizer)
+ except Exception as e:
+ logger.warning(f"Failed to read final token metadata for req {self.request_id}: {e}")
+ return
+
+ # 阶段 4:按需合并进 HTTP 输出 metadata。
+ if need_prompt_logprobs:
+ metadata["prompt_logprobs"] = meta["prompt_logprobs"]
+ metadata["prompt_token_ids"] = meta["prompt_token_ids"]
+ if meta.get("routed_experts") is not None:
+ metadata["routed_experts"] = meta["routed_experts"]
+ return
+
def get_prompt_ids(self):
return self.shm_prompt_ids.arr[: self.input_len].tolist()
@@ -297,9 +373,8 @@ def can_release(self):
ref_count_ok = self.ref_count == 1
can_released_mark = self.can_released_mark
- if self.is_aborted and can_released_mark and ref_count_ok:
- return True
-
+ # if self.is_aborted and can_released_mark and ref_count_ok:
+ # return True
ok_finished_gen_req = self.finish_status.is_finished() or self.stop_str_matched
if ok_finished_gen_req and can_released_mark and ref_count_ok and self.out_tokens_queue.is_empty():
@@ -319,23 +394,18 @@ def get_decode_need_tokens(self):
def get_first_router_need_tokens(self):
raise NotImplementedError("Subclasses should implement this method")
- def get_all_prompt_metadata(self):
- """
- return_all_prompt_logprobs mode use to return all logprobs cacul ppl
- """
- if hasattr(self, "_cache_prompt_metadata"):
- return self._cache_prompt_metadata
- metadata = {}
- cur_ids = self.shm_prompt_ids.arr[0 : self.input_len]
- all_prompts = []
- for index in range(len(cur_ids) - 1):
- tmp_dict = {int(cur_ids[index + 1]): float(self.shm_logprobs.arr[index + 1])}
- all_prompts.append([int(cur_ids[index]), tmp_dict])
-
- metadata["prompt_logprobs"] = all_prompts
- metadata["prompt_token_ids"] = [int(e) for e in cur_ids]
- self._cache_prompt_metadata = metadata
- return metadata
+ def get_output_logprobs_metadata(self, src_index: int, tokenizer=None):
+ token_id = int(self.shm_prompt_ids.arr[src_index])
+ rank = int(self.shm_logprobs.arr["rank"][src_index])
+ rank = None if rank < 0 else rank
+ return {
+ token_id: logprob_info(
+ tokenizer,
+ token_id,
+ self.shm_logprobs.arr["logprob"][src_index],
+ rank,
+ )
+ }
def is_infer_decode(self) -> bool:
"""
diff --git a/lightllm/server/core/objs/sampling_params.py b/lightllm/server/core/objs/sampling_params.py
index c39559f5f6..c122884ea1 100644
--- a/lightllm/server/core/objs/sampling_params.py
+++ b/lightllm/server/core/objs/sampling_params.py
@@ -3,6 +3,7 @@
from typing import Optional, List, Tuple, Union
from transformers import GenerationConfig
from lightllm.server.req_id_generator import MAX_BEST_OF
+from lightllm.utils.envs_utils import get_env_start_args
from .pd_kv_trans_params import PDKVTransParamObj
_SAMPLING_EPS = 1e-5
@@ -18,6 +19,7 @@
GRAMMAR_CONSTRAINT_MAX_LENGTH = int(os.getenv("LIGHTLLM_GRAMMAR_CONSTRAINT_MAX_LENGTH", 2048))
JSON_SCHEMA_MAX_LENGTH = int(os.getenv("LIGHTLLM_JSON_SCHEMA_MAX_LENGTH", 2048))
INVALID_TOKEN_IDS_MAX_LENGTH = int(os.getenv("LIGHTLLM_INVALID_TOKEN_IDS_MAX_LENGTH", 10))
+MAX_PROMPT_LOGPROBS = int(os.getenv("LIGHTLLM_MAX_PROMPT_LOGPROBS", 1024))
class StopSequence(ctypes.Structure):
@@ -305,6 +307,8 @@ class SamplingParams(ctypes.Structure):
("print_eos_token", ctypes.c_bool), # eos_id will be always ignored except the value is set to True
("disable_prompt_cache", ctypes.c_bool), # whether to disable prompt cache
("seed", ctypes.c_int64), # random seed
+ # -1 disables prompt logprobs; K >= 0 returns only the top-K prompt tokens.
+ ("prompt_logprobs", ctypes.c_int),
]
_do_sample: bool = False
@@ -341,6 +345,8 @@ def init(self, tokenizer, **kwargs):
self.add_spaces_between_special_tokens = kwargs.get("add_spaces_between_special_tokens", True)
self.print_eos_token = kwargs.get("print_eos_token", False)
self.seed = kwargs.get("seed", -1)
+ prompt_logprobs = kwargs.get("prompt_logprobs", None)
+ self.prompt_logprobs = -1 if prompt_logprobs is None else int(prompt_logprobs)
self.exponential_decay_length_penalty = ExponentialDecayLengthPenalty()
self.exponential_decay_length_penalty.initialize(kwargs.get("exponential_decay_length_penalty", (1, 1.0)))
@@ -396,15 +402,18 @@ def init(self, tokenizer, **kwargs):
def load_generation_cfg(cls, weight_dir):
try:
generation_cfg = GenerationConfig.from_pretrained(weight_dir, trust_remote_code=True).to_dict()
- cls._do_sample = generation_cfg.get("do_sample", False)
- cls._presence_penalty = generation_cfg.get("presence_penalty", 0.0)
- cls._frequency_penalty = generation_cfg.get("frequency_penalty", 0.0)
- cls._repetition_penalty = generation_cfg.get("repetition_penalty", 1.0)
- if cls._repetition_penalty is None:
- cls._repetition_penalty = 1.0
- cls._temperature = generation_cfg.get("temperature", 1.0)
- cls._top_p = generation_cfg.get("top_p", 1.0)
- cls._top_k = generation_cfg.get("top_k", -1)
+
+ def _cfg(key, default):
+ v = generation_cfg.get(key)
+ return v if v is not None else default
+
+ cls._do_sample = _cfg("do_sample", False)
+ cls._presence_penalty = _cfg("presence_penalty", 0.0)
+ cls._frequency_penalty = _cfg("frequency_penalty", 0.0)
+ cls._repetition_penalty = _cfg("repetition_penalty", 1.0)
+ cls._temperature = _cfg("temperature", 1.0)
+ cls._top_p = _cfg("top_p", 1.0)
+ cls._top_k = _cfg("top_k", -1)
except:
pass
@@ -435,6 +444,10 @@ def verify(self):
raise ValueError(
f"min_new_tokens must <= max_new_tokens, but got min {self.min_new_tokens}, max {self.max_new_tokens}."
)
+ if self.prompt_logprobs < -1 or self.prompt_logprobs > MAX_PROMPT_LOGPROBS:
+ raise ValueError(f"prompt_logprobs must be in [-1, {MAX_PROMPT_LOGPROBS}], got {self.prompt_logprobs}")
+ if self.prompt_logprobs >= 0 and not get_env_start_args().enable_prompt_logprobs:
+ raise ValueError("prompt_logprobs requires --enable_prompt_logprobs")
self._verify_allowed_token_ids()
self._verify_grammar_constraint()
@@ -487,6 +500,7 @@ def to_dict(self):
"print_eos_token": self.print_eos_token,
"disable_prompt_cache": self.disable_prompt_cache,
"seed": self.seed,
+ "prompt_logprobs": self.prompt_logprobs,
}
def to_origin_dict(self):
diff --git a/lightllm/server/core/objs/shm_array.py b/lightllm/server/core/objs/shm_array.py
index c5ad512c6b..74d64b6c5e 100644
--- a/lightllm/server/core/objs/shm_array.py
+++ b/lightllm/server/core/objs/shm_array.py
@@ -26,6 +26,13 @@ def link_shm(self):
self.arr = np.ndarray(self.shape, dtype=self.dtype, buffer=self.shm.buf)
return
+ def detach_shm(self):
+ """Close handle without unlinking (SHM persists for reuse)."""
+ if self.shm is not None:
+ self.shm.close()
+ self.shm = None
+ self.arr = None
+
def close_shm(self):
if self.shm is not None:
self.shm.close()
diff --git a/lightllm/server/core/objs/start_args_type.py b/lightllm/server/core/objs/start_args_type.py
index bfc03cd542..9e92b02e1b 100644
--- a/lightllm/server/core/objs/start_args_type.py
+++ b/lightllm/server/core/objs/start_args_type.py
@@ -1,7 +1,7 @@
from dataclasses import dataclass, field
from typing import List, Optional, Tuple
-# 只是为了更好的编程提示
+# 服务启动参数
@dataclass
@@ -10,31 +10,49 @@ class StartArgs:
default="normal",
metadata={"choices": ["normal", "pd_master", "prefill", "decode", "config_server", "visual_only"]},
)
+ performance_mode: str = field(default=None, metadata={"choices": ["personal"]})
host: str = field(default="127.0.0.1")
port: int = field(default=8000)
+ httpserver_workers: int = field(default=1)
zmq_mode: str = field(
default="ipc:///tmp/",
metadata={"help": "use socket mode or ipc mode, only can be set in ['tcp://', 'ipc:///tmp/']"},
)
- pd_master_ip: str = field(default="127.0.0.1")
+ pd_master_ip: str = field(default="0.0.0.0")
pd_master_port: int = field(default=1212)
config_server_host: str = field(default=None)
config_server_port: int = field(default=None)
config_server_visual_redis_port: int = field(default=None)
afs_image_embed_dir: str = field(default=None)
afs_embed_capacity: int = field(default=250000)
- select_p_d_node_strategy: str = field(default=None)
+ select_p_d_node_strategy: str = field(
+ default="round_robin", metadata={"choices": ["random", "round_robin", "adaptive_load"]}
+ )
model_name: str = field(default="default_model_name")
+ model_owner: Optional[str] = field(default=None)
model_dir: Optional[str] = field(default=None)
- tokenizer_mode: str = field(default="slow")
+ tokenizer_mode: str = field(default="fast")
load_way: str = field(default="HF")
max_total_token_num: Optional[int] = field(default=None)
mem_fraction: float = field(default=0.8)
batch_max_tokens: Optional[int] = field(default=None)
- eos_id: List[int] = field(default_factory=list)
+ eos_id: Optional[List[int]] = field(default=None)
tool_call_parser: Optional[str] = field(
default=None,
- metadata={"choices": ["llama3", "qwen25", "mistral", "deepseekv3", "kimi_k2", "qwen", "qwen3_coder"]},
+ metadata={
+ "choices": [
+ "qwen25",
+ "llama3",
+ "mistral",
+ "deepseekv3",
+ "qwen",
+ "deepseekv31",
+ "deepseekv32",
+ "glm47",
+ "kimi_k2",
+ "qwen3_coder",
+ ]
+ },
)
reasoning_parser: Optional[str] = field(
default=None,
@@ -53,11 +71,12 @@ class StartArgs:
"step3",
"nano_v3",
"interns1",
+ "gemma4",
]
},
)
chat_template: Optional[str] = field(default=None)
- running_max_req_size: int = field(default=512)
+ running_max_req_size: int = field(default=256)
tp: int = field(default=1)
dp: int = field(default=1)
nnodes: int = field(default=1)
@@ -69,9 +88,7 @@ class StartArgs:
use_config_server_to_init_nccl: bool = field(default=False)
trust_remote_code: bool = field(default=False)
detail_log: bool = field(default=False)
- disable_log_stats: bool = field(default=False)
- log_stats_interval: int = field(default=10)
- router_token_ratio: float = field(default=0.0)
+ router_token_ratio: float = field(default=None)
router_max_wait_tokens: int = field(default=1)
disable_aggressive_schedule: bool = field(default=False)
enable_prefill_decode_mixed: bool = field(default=False)
@@ -80,7 +97,7 @@ class StartArgs:
disable_chunked_prefill: bool = field(default=False)
diverse_mode: bool = field(default=False)
token_healing_mode: bool = field(default=False)
- output_constraint_mode: str = field(default="none", metadata={"choices": ["none", "simple", "xgrammar"]})
+ output_constraint_mode: str = field(default="none", metadata={"choices": ["outlines", "xgrammar", "none"]})
first_token_constraint_mode: bool = field(default=False)
enable_multimodal: bool = field(default=False)
disable_vision: Optional[bool] = field(default=None)
@@ -95,11 +112,12 @@ class StartArgs:
cache_capacity: int = field(default=200)
max_image_token_count: int = field(default=8192)
max_image_pixels: int = field(default=8294400)
+ disable_image_resize: bool = field(default=False)
embed_cache_storage_size: float = field(default=4)
data_type: Optional[str] = field(
default=None, metadata={"choices": ["fp16", "float16", "bf16", "bfloat16", "fp32", "float32"]}
)
- return_all_prompt_logprobs: bool = field(default=False)
+ enable_prompt_logprobs: bool = field(default=False)
use_reward_model: bool = field(default=False)
use_tgi_api: bool = field(default=False)
health_monitor: bool = field(default=False)
@@ -109,20 +127,18 @@ class StartArgs:
)
metric_gateway: Optional[str] = field(default=None)
job_name: str = field(default="lightllm")
- grouping_key: List[str] = field(default_factory=list)
+ grouping_key: List[str] = field(default_factory=lambda: [])
push_interval: int = field(default=10)
visual_node_id: int = field(default=None)
visual_infer_batch_size: int = field(default=None)
visual_send_batch_size: int = field(default=1)
- visual_gpu_ids: List[int] = field(default_factory=lambda: [0])
+ visual_gpu_ids: List[int] = field(default=None)
visual_tp: int = field(default=1)
visual_dp: int = field(default=1)
- visual_nccl_ports: List[int] = field(default=None)
visual_rpyc_port: Optional[int] = field(default=None)
audio_gpu_ids: Optional[List[int]] = field(default=None)
audio_tp: int = field(default=1)
audio_dp: int = field(default=1)
- audio_nccl_ports: Optional[List[int]] = field(default=None)
audio_infer_batch_size: Optional[int] = field(default=None)
enable_monitor_auth: bool = field(default=False)
disable_cudagraph: bool = field(default=False)
@@ -132,19 +148,19 @@ class StartArgs:
graph_split_batch_size: int = field(default=32)
graph_grow_step_size: int = field(default=16)
graph_max_len_in_batch: int = field(default=0)
- quant_type: Optional[str] = field(default=None)
+ quant_type: Optional[str] = field(default="none")
quant_cfg: Optional[str] = field(default=None)
- expert_dtype: Optional[str] = field(default=None, metadata={"choices": ["fp8", "fp4"]})
- vit_quant_type: Optional[str] = field(default=None)
+ vit_quant_type: Optional[str] = field(default="none")
vit_quant_cfg: Optional[str] = field(default=None)
+ expert_dtype: Optional[str] = field(default=None, metadata={"choices": ["fp8", "fp4"]})
llm_prefill_att_backend: List[str] = field(
- default=("auto",), metadata={"choices": ["auto", "triton", "fa3", "flashinfer"]}
+ default_factory=lambda: ["auto"], metadata={"choices": ["auto", "triton", "fa3", "flashinfer"]}
)
llm_decode_att_backend: List[str] = field(
- default=("auto",), metadata={"choices": ["auto", "triton", "fa3", "flashinfer"]}
+ default_factory=lambda: ["auto"], metadata={"choices": ["auto", "triton", "fa3", "flashinfer"]}
)
vit_att_backend: List[str] = field(
- default=("auto",), metadata={"choices": ["auto", "triton", "fa3", "sdpa", "xformers"]}
+ default_factory=lambda: ["auto"], metadata={"choices": ["auto", "triton", "fa3", "sdpa", "xformers"]}
)
llm_kv_type: str = field(
default="None", metadata={"choices": ["None", "int8kv", "int4kv", "fp8kv_sph", "fp8kv_spt", "fp8kv_dsa"]}
@@ -166,8 +182,6 @@ class StartArgs:
"eagle_with_att",
"vanilla_no_att",
"eagle_no_att",
- "qwen3next_vanilla",
- "qwen3next_eagle",
None,
]
},
@@ -180,22 +194,31 @@ class StartArgs:
pd_node_id: int = field(default=-1)
enable_cpu_cache: bool = field(default=False)
cpu_cache_storage_size: float = field(default=2)
- cpu_cache_token_page_size: int = field(default=64)
+ cpu_cache_token_page_size: int = field(default=256)
enable_disk_cache: bool = field(default=False)
disk_cache_storage_size: float = field(default=10)
disk_cache_dir: Optional[str] = field(default=None)
enable_dp_prompt_cache_fetch: bool = field(default=False)
- # zmp ports
- router_port: int = field(default=None)
- router_profiler_port: int = field(default=None)
- detokenization_port: int = field(default=None)
- http_server_port: int = field(default=None)
- visual_port: int = field(default=None)
- audio_port: int = field(default=None)
- cache_port: int = field(default=None)
- metric_port: int = field(default=None)
+ # multi-node ports (user_set; dynamic zmq ports live in ShmPortArgs)
multinode_httpmanager_port: int = field(default=12345)
- multi_level_kv_cache_port: int = field(default=None)
+
+ disable_shm_warning: bool = field(default=False)
+ dp_balancer: str = field(default="bs_balancer", metadata={"choices": ["round_robin", "bs_balancer"]})
+ enable_fused_shared_experts: bool = field(default=False)
+ enable_mps: bool = field(default=False)
+ multinode_router_gloo_port: int = field(default=20001)
+ schedule_time_interval: float = field(default=0.03)
+ use_dynamic_prompt_cache: bool = field(default=False)
+ enable_rl: bool = field(default=False)
+ enable_torch_memory_saver: bool = field(default=False)
+ enable_weight_cpu_backup: bool = field(default=False)
+ hardware_platform: str = field(default="cuda", metadata={"choices": ["cuda", "musa"]})
+ enable_torch_fallback: bool = field(default=False)
+ enable_triton_fallback: bool = field(default=False)
+
+ enable_return_routed_experts: bool = field(default=False)
+
+ weight_version: str = "default"
# hybrid attention model (Qwen3Next)
linear_att_hash_page_size: int = field(default=512)
diff --git a/lightllm/server/core/objs/token_metadata.py b/lightllm/server/core/objs/token_metadata.py
new file mode 100644
index 0000000000..11427eb76a
--- /dev/null
+++ b/lightllm/server/core/objs/token_metadata.py
@@ -0,0 +1,179 @@
+import base64
+import pickle
+from typing import Any, Dict, List, Optional
+
+import numpy as np
+
+from lightllm.utils.envs_utils import get_unique_server_name
+from lightllm.utils.log_utils import init_logger
+from lightllm.utils.shm_utils import create_or_link_shm
+
+from .logprob_utils import logprob_info
+
+logger = init_logger(__name__)
+
+
+class ReqFinalTokenMetadata:
+ """请求结束时一次性写出的 token 元信息(shm + pickle)。
+
+ 对外接口只有 ``save`` / ``read``。
+ """
+
+ def __init__(self, req):
+ self.req = req
+
+ def save(
+ self,
+ prompt_top_token_ids: Optional[np.ndarray] = None,
+ prompt_top_logprobs: Optional[np.ndarray] = None,
+ routed_experts: Optional[np.ndarray] = None,
+ ) -> None:
+ if prompt_top_token_ids is None and prompt_top_logprobs is None and routed_experts is None:
+ return
+
+ has_token_ids = prompt_top_token_ids is not None
+ has_logprobs = prompt_top_logprobs is not None
+ if has_token_ids != has_logprobs:
+ raise ValueError("prompt_top_token_ids and prompt_top_logprobs must be set together")
+
+ payload = {
+ "prompt_top_token_ids": prompt_top_token_ids,
+ "prompt_top_logprobs": prompt_top_logprobs,
+ "routed_experts": routed_experts,
+ }
+ blob = pickle.dumps(payload, protocol=pickle.HIGHEST_PROTOCOL)
+ shm = create_or_link_shm(self._shm_name(), len(blob), force_mode="create")
+ try:
+ shm.buf[: len(blob)] = blob
+ finally:
+ shm.close()
+
+ def read(self, tokenizer=None) -> Dict[str, Any]:
+ """读取并组装 HTTP 侧需要的 metadata。
+
+ Returns:
+ dict with keys:
+ - prompt_logprobs / prompt_token_ids(始终返回;无数据时为空占位)
+ - routed_experts(无数据时为 None)
+ """
+ packed_prompt_ids = None
+ packed_prompt_logprobs = None
+ packed_routed = None
+ shm = None
+ try:
+ shm = create_or_link_shm(self._shm_name(), -1, force_mode="link")
+ payload = pickle.loads(shm.buf)
+ packed_prompt_ids = payload.get("prompt_top_token_ids")
+ packed_prompt_logprobs = payload.get("prompt_top_logprobs")
+ packed_routed = payload.get("routed_experts")
+ except BaseException as e:
+ logger.warning(
+ f"Failed to read final token metadata shm for req "
+ f"{getattr(self.req, 'request_id', None)}, name={self._shm_name()}: {type(e).__name__}: {e}"
+ )
+ finally:
+ if shm is not None:
+ shm.close()
+ shm.unlink()
+
+ return {
+ "prompt_token_ids": [int(x) for x in self.req.shm_prompt_ids.arr[: self.req.input_len]],
+ "prompt_logprobs": self._build_prompt_logprobs_response(
+ tokenizer=tokenizer,
+ packed_prompt_ids=packed_prompt_ids,
+ packed_prompt_logprobs=packed_prompt_logprobs,
+ ),
+ "routed_experts": self._build_routed_experts_response(packed_routed),
+ }
+
+ def _build_prompt_logprobs_response(
+ self,
+ tokenizer,
+ packed_prompt_ids: Optional[np.ndarray],
+ packed_prompt_logprobs: Optional[np.ndarray],
+ ) -> List[Any]:
+ """组装 OpenAI 风格的 prompt_logprobs 列表。
+
+ 场景说明:
+ 1. ``input_len <= 1``:没有可预测的 prompt 位置(首 token 无前文),
+ 仅返回 ``[None]`` 占位。
+ 2. ``prompt_logprobs == 0``:不要求 top-k,只返回每个位置**真实命中**的
+ prompt token;其 logprob/rank 写在逐 token 的 ``shm_logprobs`` 里,
+ 不依赖本类 shm 中的 pickle 载荷。
+ 3. ``prompt_logprobs > 0``:返回每个位置的 top-k 候选。数据来自 Infer
+ 侧 ``save`` 写入的 ``prompt_top_token_ids/logprobs``;若 shm 缺失或
+ 未写入对应字段,则用空 dict 占位,保证列表长度仍为 ``input_len``。
+ """
+ req = self.req
+ topk = req.sample_params.prompt_logprobs
+ if req.input_len <= 1:
+ return [None]
+
+ if topk == 0:
+ # prompt_logprobs=0 返回每个位置真实命中的 prompt token,
+ # logprob/rank 存在逐 token 元信息里。
+ prompt_logprobs = [None]
+ for token_index in range(1, req.input_len):
+ token_id = int(req.shm_prompt_ids.arr[token_index])
+ rank = int(req.shm_logprobs.arr["rank"][token_index])
+ rank = None if rank < 0 else rank
+ prompt_logprobs.append(
+ {
+ token_id: logprob_info(
+ tokenizer,
+ token_id,
+ req.shm_logprobs.arr["logprob"][token_index],
+ rank,
+ )
+ }
+ )
+ return prompt_logprobs
+
+ if topk > 0:
+ prompt_logprobs = [None]
+ if packed_prompt_ids is None or packed_prompt_logprobs is None:
+ prompt_logprobs.extend({} for _ in range(req.input_len - 1))
+ return prompt_logprobs
+
+ rows = min(req.input_len - 1, packed_prompt_ids.shape[0])
+ use_topk = min(topk, packed_prompt_ids.shape[1])
+ for row_index in range(rows):
+ position_logprobs = {}
+ for index in range(use_topk):
+ top_token_id = int(packed_prompt_ids[row_index, index])
+ if top_token_id >= 0:
+ position_logprobs[top_token_id] = logprob_info(
+ tokenizer,
+ top_token_id,
+ packed_prompt_logprobs[row_index, index],
+ index + 1,
+ )
+ prompt_logprobs.append(position_logprobs)
+ if rows < req.input_len - 1:
+ prompt_logprobs.extend({} for _ in range(req.input_len - 1 - rows))
+ return prompt_logprobs
+
+ # prompt_logprobs < 0:请求未开启该字段,仍给最小占位,避免调用方 KeyError。
+ return [None]
+
+ def _build_routed_experts_response(self, packed_routed: Optional[np.ndarray]) -> Optional[Dict[str, Any]]:
+ """组装 HTTP 响应中的 routed_experts 字段。
+
+ 场景说明:
+ 1. Infer 未开启 ``--enable_return_routed_experts``,或该请求未写入
+ routing 数据时,``packed_routed`` 为 None,直接返回 None。
+ 2. 有数据时,将 ndarray 编码为 ``{shape, dtype, data}``:``data`` 为
+ C-order 原始字节的 base64,供 HTTP / 客户端按 shape+dtype 还原,
+ 避免在 JSON 里展开巨大嵌套列表。
+ """
+ if packed_routed is None:
+ return None
+ return {
+ "shape": list(packed_routed.shape),
+ "dtype": str(packed_routed.dtype),
+ "data": base64.b64encode(packed_routed.tobytes()).decode("ascii"),
+ }
+
+ def _shm_name(self) -> str:
+ service_uni_name = get_unique_server_name()
+ return f"{service_uni_name}_shm_final_token_metadata_{self.req.index_in_shm_mem}"
diff --git a/lightllm/server/detokenization/decode_req.py b/lightllm/server/detokenization/decode_req.py
index 9aa3a8effc..c77379986c 100644
--- a/lightllm/server/detokenization/decode_req.py
+++ b/lightllm/server/detokenization/decode_req.py
@@ -62,11 +62,7 @@ def stop_sequences_str_match(self) -> bool:
return False
def need_detoken(self):
- if (
- (not self.req.is_aborted)
- and (not self.req.stop_str_matched)
- and len(self.output_ids) < self.req.candetoken_out_len
- ):
+ if (not self.req.stop_str_matched) and len(self.output_ids) < self.req.candetoken_out_len:
return True
return False
@@ -83,8 +79,6 @@ def get_decode_tokens(self):
return prefix_tokens, read_tokens
def can_set_release_mark(self):
- if self.req.is_aborted:
- return True
if self.req.stop_str_matched:
return True
if (
diff --git a/lightllm/server/detokenization/manager.py b/lightllm/server/detokenization/manager.py
index 8c213914c7..58f9324859 100644
--- a/lightllm/server/detokenization/manager.py
+++ b/lightllm/server/detokenization/manager.py
@@ -17,6 +17,7 @@
import time
from lightllm.utils.log_utils import init_logger
from lightllm.utils.envs_utils import get_unique_server_name
+from lightllm.utils.shm_port_args import get_shm_port_args
logger = init_logger(__name__)
@@ -27,12 +28,13 @@ def __init__(
args: StartArgs,
):
self.args = args
+ ports = get_shm_port_args()
context = zmq.Context(2)
self.zmq_recv_socket = context.socket(zmq.PULL)
- self.zmq_recv_socket.bind(f"{args.zmq_mode}127.0.0.1:{args.detokenization_port}")
+ self.zmq_recv_socket.bind(f"{args.zmq_mode}127.0.0.1:{ports.detokenization_port}")
self.pub_to_httpserver = context.socket(zmq.PUB)
- self.pub_to_httpserver.bind(f"{args.zmq_mode}127.0.0.1:{args.http_server_port}")
+ self.pub_to_httpserver.bind(f"{args.zmq_mode}127.0.0.1:{ports.http_server_port}")
logger.info(f"pub_to_httpserver sendhwm {self.pub_to_httpserver.getsockopt(zmq.SNDHWM)}")
self.tokenizer = get_tokenizer(args.model_dir, args.tokenizer_mode, trust_remote_code=args.trust_remote_code)
self.all_special_ids = set(self.tokenizer.all_special_ids)
diff --git a/lightllm/server/embed_cache/impl/naive_memory_cache.py b/lightllm/server/embed_cache/impl/naive_memory_cache.py
index 0dad890e2c..878542ddd8 100644
--- a/lightllm/server/embed_cache/impl/naive_memory_cache.py
+++ b/lightllm/server/embed_cache/impl/naive_memory_cache.py
@@ -59,7 +59,11 @@ def _check_and_set_new_id_range(self, alloced_token_num):
else:
while True:
try:
- config_server_ip_port = f"{self.args.config_server_host}:{self.args.config_server_port}"
+ from lightllm.utils.shm_port_args import get_shm_port_args
+
+ config_server_ip_port = (
+ f"{self.args.config_server_host}:{get_shm_port_args().config_server_port}"
+ )
url = f"http://{config_server_ip_port}/allocate_global_unique_multimodal_id_range"
response = requests.get(url)
if response.status_code == 200:
diff --git a/lightllm/server/embed_cache/manager.py b/lightllm/server/embed_cache/manager.py
index 5de4df4ab3..a6d5d6788c 100644
--- a/lightllm/server/embed_cache/manager.py
+++ b/lightllm/server/embed_cache/manager.py
@@ -8,6 +8,7 @@
from lightllm.server.embed_cache.impl.naive_memory_cache import InMemoryCache
from rpyc.utils.classic import obtain
from lightllm.utils.envs_utils import get_unique_server_name
+from lightllm.utils.shm_port_args import get_shm_port_args
class CacheServer(rpyc.Service):
@@ -62,7 +63,7 @@ def start_cache_manager(args: StartArgs, pipe_writer):
from rpyc.utils.server import ThreadedServer
import lightllm.utils.rpyc_fix_utils as _
- t = ThreadedServer(service, port=args.cache_port, protocol_config={"allow_pickle": True})
+ t = ThreadedServer(service, port=get_shm_port_args().cache_port, protocol_config={"allow_pickle": True})
pipe_writer.send("init ok")
t.start()
diff --git a/lightllm/server/health_monitor/manager.py b/lightllm/server/health_monitor/manager.py
index c7373494b1..f6ece99ce6 100644
--- a/lightllm/server/health_monitor/manager.py
+++ b/lightllm/server/health_monitor/manager.py
@@ -12,6 +12,7 @@
from lightllm.utils.log_utils import init_logger
from lightllm.utils.graceful_utils import graceful_registry
from lightllm.utils.envs_utils import get_unique_server_name
+from lightllm.utils.shm_port_args import get_shm_port_args
logger = init_logger(__name__)
@@ -97,7 +98,7 @@ def start_health_check_process(args, pipe_writer):
logger.info(f"health monitor care process ids {all_process_ids}")
global consecutive_failures
- host, port = args.host, args.port
+ host, port = args.host, get_shm_port_args().port
url = f"http://{host}:{port}/health".format(host=host, port=port)
interval_seconds = int(os.environ.get("HEALTH_CHECK_INTERVAL_SECONDS", 88))
logger.info("Waiting for the server to start up.")
diff --git a/lightllm/server/httpserver/manager.py b/lightllm/server/httpserver/manager.py
index 0f1b873111..9a316160d7 100644
--- a/lightllm/server/httpserver/manager.py
+++ b/lightllm/server/httpserver/manager.py
@@ -31,24 +31,28 @@
from lightllm.server.router.dynamic_prompt.shared_arr import SharedInt
from lightllm.utils.log_utils import init_logger
from lightllm.server.metrics.manager import MetricClient
+from .rl_controller import HttpRlController
+from .manager_ext import HttpRlManagerHelper
from lightllm.utils.statics_utils import MovingAverage
from lightllm.utils.config_utils import get_vocab_size
from lightllm.utils.envs_utils import get_unique_server_name
+from lightllm.utils.shm_port_args import get_shm_port_args
from lightllm.utils.error_utils import ClientDisconnected, PDPrefillNodeStopGenToken
from rpyc.utils.classic import obtain
logger = init_logger(__name__)
-class HttpServerManager:
+class HttpServerManager(HttpRlManagerHelper, object):
def __init__(
self,
args: StartArgs,
):
self.args: StartArgs = args
+ ports = get_shm_port_args()
context = zmq.asyncio.Context(2)
self.send_to_router = context.socket(zmq.PUSH)
- self.send_to_router.connect(f"{args.zmq_mode}127.0.0.1:{args.router_port}")
+ self.send_to_router.connect(f"{args.zmq_mode}127.0.0.1:{ports.router_port}")
self.multinode_req_manager = None
self.nnodes = args.nnodes
@@ -56,7 +60,6 @@ def __init__(
self._resource_lock = AsyncLock(self._shm_lock_pool.get_lock_context(0))
self._run_reqs_count_lock = AsyncLock(self._shm_lock_pool.get_lock_context(1))
self.node_rank = args.node_rank
- self.disable_abort = args.nnodes > 1 and args.dp == 1 # mulitnode dp=1 mode, disable abort
self.is_multinode_tp = args.dp == 1 and args.nnodes > 1
self.is_multinode_tp_master = args.dp == 1 and args.nnodes > 1 and args.node_rank == 0
self.is_multinode_tp_slave = args.dp == 1 and args.nnodes > 1 and args.node_rank > 0
@@ -66,41 +69,41 @@ def __init__(
for child_ip in args.child_ips:
context = zmq.asyncio.Context(2)
self.multinode_req_manager.append(context.socket(zmq.PUSH))
- self.multinode_req_manager[-1].connect(f"tcp://{child_ip}:{args.multinode_httpmanager_port}")
+ self.multinode_req_manager[-1].connect(f"tcp://{child_ip}:{ports.multinode_httpmanager_port}")
logger.info(
- f"HttpServerManager connected to child node at {child_ip}:{args.multinode_httpmanager_port}"
+ f"HttpServerManager connected to child node at {child_ip}:{ports.multinode_httpmanager_port}"
)
else:
context = zmq.asyncio.Context(2)
self.multinode_req_manager = context.socket(zmq.PULL)
- self.multinode_req_manager.bind(f"tcp://*:{args.multinode_httpmanager_port}")
+ self.multinode_req_manager.bind(f"tcp://*:{ports.multinode_httpmanager_port}")
logger.info(
- f"HttpServerManager listening for child node requests on *:{args.multinode_httpmanager_port}"
+ f"HttpServerManager listening for master node requests on *:{ports.multinode_httpmanager_port}"
)
self.enable_multimodal = args.enable_multimodal
if self.enable_multimodal:
- self.cache_client = rpyc.connect("localhost", args.cache_port, config={"allow_pickle": True})
+ self.cache_client = rpyc.connect("localhost", ports.cache_port, config={"allow_pickle": True})
self.cache_client._channel.stream.sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1)
if not self.args.disable_vision:
self.send_to_visual = context.socket(zmq.PUSH)
- self.send_to_visual.connect(f"{args.zmq_mode}127.0.0.1:{args.visual_port}")
+ self.send_to_visual.connect(f"{args.zmq_mode}127.0.0.1:{ports.visual_port}")
if not self.args.disable_audio:
self.send_to_audio = context.socket(zmq.PUSH)
- self.send_to_audio.connect(f"{args.zmq_mode}127.0.0.1:{args.audio_port}")
+ self.send_to_audio.connect(f"{args.zmq_mode}127.0.0.1:{ports.audio_port}")
if args.enable_cpu_cache and not self.args.enable_multimodal:
self.send_to_multi_level_kv_cache = context.socket(zmq.PUSH)
- self.send_to_multi_level_kv_cache.connect(f"{args.zmq_mode}127.0.0.1:{args.multi_level_kv_cache_port}")
+ self.send_to_multi_level_kv_cache.connect(f"{args.zmq_mode}127.0.0.1:{ports.multi_level_kv_cache_port}")
self.shm_req_manager = ShmReqManager()
# recv from detokenization
self.zmq_recv_socket = context.socket(zmq.SUB)
- self.zmq_recv_socket.connect(f"{args.zmq_mode}127.0.0.1:{args.http_server_port}")
+ self.zmq_recv_socket.connect(f"{args.zmq_mode}127.0.0.1:{ports.http_server_port}")
self.zmq_recv_socket.setsockopt(zmq.SUBSCRIBE, b"")
self.tokenizer = get_tokenizer(args.model_dir, args.tokenizer_mode, trust_remote_code=args.trust_remote_code)
@@ -109,7 +112,7 @@ def __init__(
self.forwarding_queue: AsyncQueue = None # p d 分离模式使用的转发队列, 需要延迟初始化
self.max_req_total_len = args.max_req_total_len
- self.metric_client = MetricClient(args.metric_port)
+ self.metric_client = MetricClient(ports.metric_port)
self.pd_mode: NodeRole = NodeRole(self.args.run_mode)
assert self.pd_mode in [NodeRole.NORMAL, NodeRole.P, NodeRole.D]
@@ -123,6 +126,8 @@ def __init__(
self.latest_success_infer_time_mark = SharedInt(f"{get_unique_server_name()}_latest_success_infer_time_mark")
self.latest_success_infer_time_mark.set_value(int(time.time()))
+ self.rl_controller: Optional[HttpRlController] = HttpRlController(self) if args.enable_rl else None
+
self.run_reqs_count_mark = SharedInt(f"{get_unique_server_name()}_run_reqs_count_mark")
self.run_reqs_count_mark.set_value(0)
@@ -335,6 +340,18 @@ async def generate(
self.run_reqs_count_mark.set_value(self.run_reqs_count_mark.get_value() + 1)
try:
+ # RL:进入 generation admission。若当前处于 pause_generation / abort,
+ # 在此阻塞等待恢复;若被标记 abort,则直接返回 FINISHED_ABORTED 空结果,
+ # 不进入后续 encode / 调度(请求尚未占用 router 资源)。
+ if self.rl_controller is not None:
+ should_abort = await self.rl_controller.wait_if_generation_paused(group_request_id)
+ if should_abort:
+ for output in self.rl_controller.build_aborted_generation_outputs(
+ group_request_id, sampling_params
+ ):
+ yield output
+ return
+
original_multimodal_params = None
if self.is_multinode_tp_master:
original_multimodal_params = copy.deepcopy(multimodal_params)
@@ -430,6 +447,10 @@ async def generate(
req_status = ReqStatus(group_request_id, multimodal_params, req_objs, start_time)
self.req_id_to_out_inf[group_request_id] = req_status
+ # RL:请求已登记到 req_id_to_out_inf 并即将转发下游,从 admission gate
+ # 注销,避免 pause 统计里仍把它算作“等待准入”的 pending 请求。
+ if self.rl_controller is not None:
+ await self.rl_controller.unregister_generation_admission(group_request_id)
await self.transfer_to_next_module_or_node(
prompt, sampling_params, original_multimodal_params, req_status.group_req_objs
@@ -467,11 +488,13 @@ async def generate(
yield sub_req_id, request_output, metadata, finish_status
- except (ClientDisconnected, Exception) as e:
- logger.warning(f"group_request_id: {group_request_id} has exception {str(e)}")
-
+ except (asyncio.CancelledError, BaseException) as e:
if isinstance(e, ClientDisconnected):
logger.warning(f"group_request_id: {group_request_id} {e.reason}")
+ elif isinstance(e, asyncio.CancelledError):
+ logger.warning(f"group_request_id: {group_request_id} has been cancelled")
+ else:
+ logger.warning(f"group_request_id: {group_request_id} has exception {str(e)}")
# error need to release multimodel resources.
# 对于还没有形成正式请求对象管理的多模态资源,需要单独自己释放
@@ -482,6 +505,10 @@ async def generate(
await self.abort(group_request_id)
raise e
finally:
+ # RL:兜底注销 generation admission(含中途异常 / abort 提前 return),
+ # 防止 pending 请求泄漏导致 pause 无法正确结束。
+ if self.rl_controller is not None:
+ await self.rl_controller.unregister_generation_admission(group_request_id)
async with self._run_reqs_count_lock:
self.run_reqs_count_mark.set_value(self.run_reqs_count_mark.get_value() - 1)
return
@@ -501,7 +528,6 @@ def _count_multimodal_tokens(self, multimodal_params: MultimodalParams) -> Tuple
return image_tokens, audio_tokens
async def _log_req_header(self, request_headers, group_request_id: int):
-
x_request_id = request_headers.get("X-Request-Id", "")
x_session_id = request_headers.get("X-Session-Id", "")
@@ -527,11 +553,7 @@ async def _encode(
f"the request is rejected before tokenization."
)
if self.enable_multimodal:
- assert (
- len(multimodal_params.images + multimodal_params.audios) <= self.args.cache_capacity
- ), "too many multimodal items!"
- if multimodal_params.audios:
- assert not self.args.disable_audio, "audio multimodal not enabled"
+ multimodal_params.verify_resource_limits()
await self._alloc_multimodal_resources(multimodal_params, sampling_params)
prompt_ids = await asyncio.to_thread(
self.tokenizer.encode,
@@ -556,12 +578,24 @@ async def _encode(
# 这里的校验对多模态不是很充分, to do
if all(isinstance(e, int) for e in prompt):
- if not self.enable_multimodal and not self.pd_mode.is_D():
+ if not self.enable_multimodal and self.pd_mode.is_P_or_NORMAL():
if all(e < self.vocab_size for e in prompt):
return prompt
else:
raise ValueError("prompt List[int] format contain id > vocab_size")
else:
+ if self.enable_multimodal and self.pd_mode.is_P_or_NORMAL():
+ multimodal_params.verify_resource_limits()
+ await self._alloc_multimodal_resources(multimodal_params, sampling_params)
+ # prompt 已是 List[int](预分词),但仍可能携带 images/audios。
+ # 多模态 tokenizer.encode 不会对 int 再做文本分词,而是在已有 ids 上
+ # 展开 image/audio 占位 token(如
-> 连续 token_id)。
+ # 因此,这里需要再次调用 tokenizer.encode 来展开 image/audio 占位 token。
+ return self.tokenizer.encode(
+ prompt,
+ multimodal_params,
+ add_special_tokens=sampling_params.add_special_tokens,
+ )
return prompt
else:
raise ValueError(f"prompt format error, get type{type(prompt)}")
@@ -579,7 +613,6 @@ async def _check_and_repair_length(self, prompt_ids: List[int], sampling_params:
real_supported_max_req_total_len = self.get_real_supported_max_req_total_len()
if prompt_tokens + sampling_params.max_new_tokens > real_supported_max_req_total_len:
-
# 修改默认逻辑,如果 prompt_tokens + max_new_tokens 长度超过总的允许长度,则将
# 修改 max_new_tokens 的值,使其满足合法约束。
new_max_new_tokens = real_supported_max_req_total_len - prompt_tokens
@@ -685,10 +718,7 @@ async def _wait_to_token_package(
except asyncio.TimeoutError:
pass
- if req_status.aborted:
- raise Exception(f"req_id {group_request_id} aborted notifyed by other module")
-
- if not self.disable_abort and request is not None and await request.is_disconnected():
+ if request is not None and await request.is_disconnected():
await self.abort(group_request_id)
raise ClientDisconnected(
group_request_id=group_request_id, reason="_wait_to_token_package check network disconnected"
@@ -807,7 +837,7 @@ def _get_router_profiler_client(self):
self.router_profiler_client = retry(max_attempts=20, wait_time=0.5)(rpyc.connect)(
"localhost",
- self.args.router_profiler_port,
+ get_shm_port_args().router_profiler_port,
config={"allow_pickle": True},
)
self.router_profiler_client._channel.stream.sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1)
@@ -821,7 +851,6 @@ async def recycle_resource_loop(self):
pre_time_mark = time.time()
while True:
-
try:
await asyncio.wait_for(self.recycle_event.wait(), timeout=0.02)
except asyncio.TimeoutError:
@@ -837,16 +866,11 @@ async def recycle_resource_loop(self):
for req_status in release_req_status:
self.req_id_to_out_inf.pop(req_status.group_req_objs.group_req_id, None)
- _is_aborted = False
for req in req_status.group_req_objs.shm_req_objs:
- _is_aborted = _is_aborted or req.is_aborted
logger.debug(f"httpserver release req_id {req.request_id}, index {req.index_in_shm_mem}")
await self.shm_req_manager.async_put_back_req_obj(req)
await self.shm_req_manager.async_release_req_index(req.index_in_shm_mem)
await self._release_multimodal_resources(req_status.group_req_objs.multimodal_params)
- if _is_aborted:
- req_status.aborted = True
- logger.debug(f"mark req_id {req_status.group_req_objs.group_req_id} aborted in recycle loop")
# 先保留这个关键得日志,用于方便定位重构中的问题。
if time.time() - pre_time_mark > 120:
@@ -898,12 +922,12 @@ async def handle_loop(self):
for _ in range(read_token_count):
if not req.out_tokens_queue.is_empty():
-
text, src_index, special, count_output_tokens = req.out_tokens_queue.peek()
- req.cumlogprob += float(req.shm_logprobs.arr[src_index])
+ token_logprob = float(req.shm_logprobs.arr["logprob"][src_index])
+ req.cumlogprob += token_logprob
metadata = {
"id": int(req.shm_prompt_ids.arr[src_index]),
- "logprob": float(req.shm_logprobs.arr[src_index]),
+ "logprob": token_logprob,
"cumlogprob": float(req.cumlogprob) / count_output_tokens,
"special": special,
"count_output_tokens": count_output_tokens,
@@ -912,26 +936,30 @@ async def handle_loop(self):
"disk_prompt_cache_len": req.disk_prompt_cache_len,
"mtp_accepted_token_num": req.mtp_accepted_token_num,
}
- if self.args.return_all_prompt_logprobs:
- metadata.update(req.get_all_prompt_metadata())
+ metadata["logprobs"] = req.get_output_logprobs_metadata(src_index, self.tokenizer)
if self.args.use_reward_model:
metadata["score"] = float(req.reward_score)
- req.out_tokens_queue.pop_no_ret()
-
finished_token_index = (
req.stop_str_matched_token_index if req.stop_str_matched else req.finish_token_index
)
if finished_token_index != src_index:
- token_list.append((req_id, text, metadata, FinishStatus()))
+ finish_status = FinishStatus()
else:
if req.stop_str_matched:
finish_status = FinishStatus(FinishStatus.FINISHED_STOP)
else:
finish_status = FinishStatus(req.finish_status.status)
- token_list.append((req_id, text, metadata, finish_status))
+ await req.merge_final_token_metadata(
+ metadata,
+ self.tokenizer,
+ enable_return_routed_experts=self.args.enable_return_routed_experts,
+ )
+
+ req.out_tokens_queue.pop_no_ret()
+ token_list.append((req_id, text, metadata, finish_status))
else:
break
@@ -957,7 +985,6 @@ def __init__(self, group_request_id, multimodal_params, req_objs: List[Req], sta
time_mark=start_time,
)
self.out_token_info_list = []
- self.aborted = False
def can_release(self):
for req in self.group_req_objs.shm_req_objs:
diff --git a/lightllm/server/httpserver/manager_ext.py b/lightllm/server/httpserver/manager_ext.py
new file mode 100644
index 0000000000..212de48e4e
--- /dev/null
+++ b/lightllm/server/httpserver/manager_ext.py
@@ -0,0 +1,60 @@
+"""HttpServerManager 的扩展 Mixin。
+
+通过多继承挂到 :class:`HttpServerManager` 上,把 RL 控制面 HTTP 转发接口
+从 manager 主体中拆出,避免 manager.py 继续膨胀。
+
+约定:宿主类在 ``--enable_rl`` 时提供非空 ``self.rl_controller``;
+RL HTTP 路由也仅在该开关下挂载,因此这些接口只应在 enable_rl 场景被调用。
+"""
+
+from typing import Tuple
+
+from lightllm.server.io_struct import (
+ AbortReq,
+ DestroyWeightsUpdateGroupReq,
+ FlushCacheReq,
+ InitWeightsUpdateGroupReq,
+ ReleaseMemoryReq,
+ ResumeMemoryReq,
+ RlOpRsp,
+ UpdateWeightsFromDistributedReq,
+ UpdateWeightsFromIPCReq,
+ UpdateWeightsFromTensorReq,
+)
+
+
+class HttpRlManagerHelper:
+ """RL 控制面接口 Mixin:一律转发到 ``self.rl_controller``。"""
+
+ async def abort_request(self, request: AbortReq) -> Tuple[bool, str]:
+ return await self.rl_controller.abort_request(request)
+
+ async def pause_generation(self):
+ return await self.rl_controller.pause_generation()
+
+ async def continue_generation(self):
+ return await self.rl_controller.continue_generation()
+
+ async def flush_cache(self, request: FlushCacheReq):
+ return await self.rl_controller.flush_cache(request)
+
+ async def release_memory_occupation(self, request: ReleaseMemoryReq):
+ return await self.rl_controller.release_memory_occupation(request)
+
+ async def resume_memory_occupation(self, request: ResumeMemoryReq):
+ return await self.rl_controller.resume_memory_occupation(request)
+
+ async def init_weights_update_group(self, request: InitWeightsUpdateGroupReq):
+ return await self.rl_controller.init_weights_update_group(request)
+
+ async def destroy_weights_update_group(self, request: DestroyWeightsUpdateGroupReq):
+ return await self.rl_controller.destroy_weights_update_group(request)
+
+ async def update_weights_from_distributed(self, request: UpdateWeightsFromDistributedReq):
+ return await self.rl_controller.update_weights_from_distributed(request)
+
+ async def update_weights_from_tensor(self, request: UpdateWeightsFromTensorReq) -> RlOpRsp:
+ return await self.rl_controller.update_weights_from_tensor(request)
+
+ async def update_weights_from_ipc(self, request: UpdateWeightsFromIPCReq) -> RlOpRsp:
+ return await self.rl_controller.update_weights_from_ipc(request)
diff --git a/lightllm/server/httpserver/pd_loop.py b/lightllm/server/httpserver/pd_loop.py
index dcf0c89fed..297a6fe370 100644
--- a/lightllm/server/httpserver/pd_loop.py
+++ b/lightllm/server/httpserver/pd_loop.py
@@ -21,6 +21,7 @@
from lightllm.server.core.objs import StartArgs
from lightllm.server.core.objs import SamplingParams
from lightllm.utils.error_utils import PDPrefillNodeStopGenToken
+from lightllm.utils.shm_port_args import get_shm_port_args
logger = init_logger(__name__)
@@ -96,7 +97,7 @@ async def _pd_handle_task(manager: HttpServerManager, pd_master_obj: PD_Master_O
# 发送注册信息
regist_json = {
"node_id": manager.args.pd_node_id,
- "client_ip_port": f"{manager.host_ip}:{manager.args.port}",
+ "client_ip_port": f"{manager.host_ip}:{get_shm_port_args().port}",
"mode": manager.pd_mode.value,
"start_args": args_dict,
}
@@ -180,11 +181,11 @@ async def _get_pd_master_objs(args: StartArgs) -> Optional[Dict[int, PD_Master_O
# node_id 为 0
if not use_config_server:
ans = dict()
- ans[0] = PD_Master_Obj(node_id=0, host_ip_port=f"{args.pd_master_ip}:{args.pd_master_port}")
+ ans[0] = PD_Master_Obj(node_id=0, host_ip_port=f"{args.pd_master_ip}:{get_shm_port_args().pd_master_port}")
return ans
# 使用 config_server 服务来发现所有的 pd_master 节点。
- uri = f"ws://{args.config_server_host}:{args.config_server_port}/registered_objects"
+ uri = f"ws://{args.config_server_host}:{get_shm_port_args().config_server_port}/registered_objects"
try:
async with httpx.AsyncClient() as client:
@@ -256,6 +257,6 @@ def _get_load_info() -> dict:
mean_node_load = sum(current_load) / len(current_load)
load_info = {
"total_token_usage_rate": mean_node_load,
- "client_ip_port": f"{g_objs.httpserver_manager.host_ip}:{g_objs.args.port}",
+ "client_ip_port": f"{g_objs.httpserver_manager.host_ip}:{get_shm_port_args().port}",
}
return load_info
diff --git a/lightllm/server/httpserver/rl_controller.py b/lightllm/server/httpserver/rl_controller.py
new file mode 100644
index 0000000000..eda841a274
--- /dev/null
+++ b/lightllm/server/httpserver/rl_controller.py
@@ -0,0 +1,301 @@
+import asyncio
+import socket
+import time
+import rpyc
+from contextlib import asynccontextmanager
+from typing import Tuple
+from rpyc.utils.classic import obtain
+from lightllm.server.core.objs import FinishStatus, SamplingParams
+from lightllm.server.io_struct import (
+ AbortReq,
+ FlushCacheReq,
+ RlOpReq,
+ RlOpRsp,
+ InitWeightsUpdateGroupReq,
+ DestroyWeightsUpdateGroupReq,
+ ReleaseMemoryReq,
+ ResumeMemoryReq,
+ UpdateWeightsFromDistributedReq,
+ UpdateWeightsFromIPCReq,
+ UpdateWeightsFromTensorReq,
+)
+from lightllm.utils.log_utils import init_logger
+from lightllm.utils.shm_port_args import get_shm_port_args
+
+logger = init_logger(__name__)
+
+
+class _GenerationPauseGate:
+ """Generation pause gate.
+
+ Requests are tracked from admission until req_id_to_out_inf takes over.
+ """
+
+ _RUNNING = 0
+ _PAUSED = 1
+
+ def __init__(self) -> None:
+ self._state = self._RUNNING
+ self._pending_request_abort_events = {}
+ self._lock = asyncio.Lock()
+ self._resume_event = asyncio.Event()
+ self._resume_event.set()
+
+ @asynccontextmanager
+ async def pause_and_abort_context(self):
+ """Enter pause mode once and tell the caller whether it owns the abort pass."""
+ async with self._lock:
+ if self._state != self._RUNNING:
+ do_abort = False
+ else:
+ # New requests after this point should wait at the gate until resume.
+ # abort_all below is responsible only for requests already present
+ # when that abort_all call takes its snapshot.
+ self._state = self._PAUSED
+ self._resume_event.clear()
+ do_abort = True
+ yield do_abort
+
+ async def unregister_pending_request(self, request_id: int):
+ async with self._lock:
+ self._pending_request_abort_events.pop(request_id, None)
+
+ async def wait_until_resumed_or_aborted(self, request_id: int) -> bool:
+ """Enter admission and return True if this request should abort."""
+ async with self._lock:
+ abort_event = asyncio.Event()
+ self._pending_request_abort_events[request_id] = abort_event
+ if self._state == self._RUNNING:
+ return False
+
+ resume_event = self._resume_event
+
+ resume_task = asyncio.create_task(resume_event.wait())
+ abort_task = asyncio.create_task(abort_event.wait())
+ try:
+ done, pending = await asyncio.wait({resume_task, abort_task}, return_when=asyncio.FIRST_COMPLETED)
+ except asyncio.CancelledError:
+ resume_task.cancel()
+ abort_task.cancel()
+ await asyncio.gather(resume_task, abort_task, return_exceptions=True)
+ await self.unregister_pending_request(request_id)
+ raise
+
+ for task in pending:
+ task.cancel()
+ if pending:
+ await asyncio.gather(*pending, return_exceptions=True)
+ if abort_task in done and abort_event.is_set():
+ await self.unregister_pending_request(request_id)
+ return True
+ return False
+
+ async def abort_pending_requests(self, request_ids) -> None:
+ async with self._lock:
+ for request_id in request_ids:
+ abort_event = self._pending_request_abort_events.get(request_id)
+ if abort_event is not None:
+ abort_event.set()
+
+ async def snapshot_and_abort_pending_requests(self):
+ async with self._lock:
+ request_ids = list(self._pending_request_abort_events.keys())
+ for request_id in request_ids:
+ self._pending_request_abort_events[request_id].set()
+ return request_ids
+
+ def has_pending_request(self, request_id: int) -> bool:
+ return request_id in self._pending_request_abort_events
+
+ async def resume(self) -> None:
+ async with self._lock:
+ self._state = self._RUNNING
+ self._resume_event.set()
+
+
+class HttpRlController:
+ def __init__(self, manager) -> None:
+ self.manager = manager
+ self.args = manager.args
+ self._generation_gate = _GenerationPauseGate()
+
+ async def wait_if_generation_paused(self, request_id: int) -> bool:
+ """Enter generation admission and return True if this request should abort."""
+ return await self._generation_gate.wait_until_resumed_or_aborted(request_id)
+
+ async def unregister_generation_admission(self, request_id: int) -> None:
+ await self._generation_gate.unregister_pending_request(request_id)
+
+ def build_aborted_generation_outputs(self, group_request_id: int, sampling_params: SamplingParams):
+ """构造被 pause/abort 拦截请求的空结果,供 HTTP generate 流正常收尾。
+
+ 场景:RL ``pause_generation`` / ``abort_request`` 期间,新请求会在
+ admission gate 处被标记 abort,此时请求尚未进入 router 队列,也没有
+ 真实生成 token。为了让 stream / non-stream API 仍能按统一协议收到
+ ``FINISHED_ABORTED`` 并结束生成循环(而不是挂起或抛错),这里合成
+ n 路空输出。
+ """
+ finish_status = FinishStatus()
+ finish_status.set_status(FinishStatus.FINISHED_ABORTED)
+ metadata = {"prompt_tokens": 0, "count_output_tokens": 0}
+ for i in range(sampling_params.n):
+ yield group_request_id + i, "", metadata, finish_status
+
+ async def _wait_for_abort_released(
+ self,
+ request_ids,
+ abort_when_running_request_ids=(),
+ timeout: float = 60.0,
+ ) -> Tuple[bool, str]:
+ """Wait until aborted work has left both the pause gate and running map.
+
+ request_ids is the fixed set of requests selected by abort_request().
+ Requests entering the pause gate later are not part of this wait.
+ """
+ start_time = time.time()
+ request_ids = set(request_ids)
+ abort_when_running_request_ids = set(abort_when_running_request_ids)
+ while True:
+ has_unreleased_req = False
+ for group_req_id in request_ids:
+ if self._generation_gate.has_pending_request(group_req_id):
+ has_unreleased_req = True
+ break
+
+ if group_req_id in self.manager.req_id_to_out_inf:
+ if group_req_id in abort_when_running_request_ids:
+ await self.manager.abort(group_req_id)
+ abort_when_running_request_ids.remove(group_req_id)
+ has_unreleased_req = True
+ break
+
+ if not has_unreleased_req:
+ return True, ""
+
+ if time.time() - start_time > timeout:
+ error_msg = (
+ f"abort request wait release timeout, request_ids_count={len(request_ids)}, timeout={timeout}s"
+ )
+ logger.error(error_msg)
+ return False, error_msg
+
+ await asyncio.sleep(0.02)
+ return True, ""
+
+ async def abort_request(self, request: AbortReq) -> Tuple[bool, str]:
+ """Abort one request, or snapshot and abort all requests present now."""
+ request_id = request.request_id
+ if request.abort_all:
+ # Snapshot before issuing aborts: this abort_all clears current
+ # pending/running work, while future paused requests keep waiting.
+ pending_request_ids = set(await self._generation_gate.snapshot_and_abort_pending_requests())
+ running_request_ids = set(self.manager.req_id_to_out_inf.keys())
+ abort_request_ids = pending_request_ids | running_request_ids
+ for group_req_id in running_request_ids:
+ await self.manager.abort(group_req_id)
+ return await self._wait_for_abort_released(
+ request_ids=abort_request_ids, abort_when_running_request_ids=pending_request_ids
+ )
+
+ if request_id is None:
+ return True, ""
+
+ await self._generation_gate.abort_pending_requests((request_id,))
+ await self.manager.abort(request_id)
+ return await self._wait_for_abort_released(request_ids=(request_id,))
+
+ async def pause_generation(self):
+ """Pause future generations and drain only the requests already present."""
+ async with self._generation_gate.pause_and_abort_context() as do_abort:
+ if not do_abort:
+ return
+ while True:
+ success, msg = await self.abort_request(AbortReq(request_id=None, abort_all=True))
+ if success:
+ break
+ logger.warning(f"pause_generation abort_all still waiting: {msg}")
+ await asyncio.sleep(1.0)
+
+ async def continue_generation(self):
+ """Release requests waiting at the pause gate."""
+ await self._generation_gate.resume()
+
+ def _call_rl_op_sync(self, req: RlOpReq) -> RlOpRsp:
+ from lightllm.utils.retry_utils import retry
+
+ conn = retry(max_attempts=20, wait_time=0.5)(rpyc.connect)(
+ "localhost",
+ get_shm_port_args().rl_rpyc_port,
+ config={"allow_pickle": True, "sync_request_timeout": 600},
+ )
+ try:
+ conn._channel.stream.sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1)
+ return obtain(conn.root.rl_op(req))
+ finally:
+ try:
+ conn.close()
+ except BaseException:
+ pass
+
+ async def _call_rl_op(self, op_name: str, op_args=None) -> RlOpRsp:
+ req = RlOpReq(op_name=op_name, op_args=op_args)
+ try:
+ return await asyncio.to_thread(self._call_rl_op_sync, req)
+ except BaseException as e:
+ logger.exception(f"rl op {op_name} failed: {e}")
+ return RlOpRsp(
+ success=False,
+ msg=f"rl op {op_name} error: {e}",
+ op_name=op_name,
+ )
+
+ async def flush_cache(self, request: FlushCacheReq):
+ return await self._call_rl_op("flush_cache", request)
+
+ async def release_memory_occupation(self, request: ReleaseMemoryReq):
+ assert (
+ len(self.manager.req_id_to_out_inf) == 0
+ ), "there are still requests running, cannot release memory occupation"
+ return await self._call_rl_op("release_memory_occupation", request.tags)
+
+ async def resume_memory_occupation(self, request: ResumeMemoryReq):
+ return await self._call_rl_op("resume_memory_occupation", request.tags)
+
+ async def init_weights_update_group(self, request: InitWeightsUpdateGroupReq):
+ return await self._call_rl_op("init_weights_update_group", request)
+
+ async def destroy_weights_update_group(self, request: DestroyWeightsUpdateGroupReq):
+ return await self._call_rl_op("destroy_weights_update_group", request)
+
+ async def update_weights_from_distributed(self, request: UpdateWeightsFromDistributedReq):
+ if request.abort_all_requests:
+ success, msg = await self.abort_request(AbortReq(abort_all=True))
+ if not success:
+ return RlOpRsp(success=False, msg=msg, op_name="update_weights_from_distributed")
+ if request.flush_cache:
+ ret = await self.flush_cache(FlushCacheReq())
+ if not ret.success:
+ return ret
+ return await self._call_rl_op("update_weights_from_distributed", request)
+
+ async def update_weights_from_tensor(self, request: UpdateWeightsFromTensorReq) -> RlOpRsp:
+ if request.abort_all_requests:
+ success, msg = await self.abort_request(AbortReq(abort_all=True))
+ if not success:
+ return RlOpRsp(success=False, msg=msg, op_name="update_weights_from_tensor")
+ if request.flush_cache:
+ ret = await self.flush_cache(FlushCacheReq())
+ if not ret.success:
+ return ret
+ return await self._call_rl_op("update_weights_from_tensor", request)
+
+ async def update_weights_from_ipc(self, request: UpdateWeightsFromIPCReq) -> RlOpRsp:
+ if request.abort_all_requests:
+ success, msg = await self.abort_request(AbortReq(abort_all=True))
+ if not success:
+ return RlOpRsp(success=False, msg=msg, op_name="update_weights_from_ipc")
+ if request.flush_cache:
+ ret = await self.flush_cache(FlushCacheReq())
+ if not ret.success:
+ return ret
+ return await self._call_rl_op("update_weights_from_ipc", request)
diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py
index 104da9f26e..54ff7231b6 100644
--- a/lightllm/server/httpserver_for_pd_master/manager.py
+++ b/lightllm/server/httpserver_for_pd_master/manager.py
@@ -21,6 +21,7 @@
from lightllm.server.httpserver.manager import AsyncQueue
from lightllm.utils.error_utils import ClientDisconnected, ServerBusyError
from lightllm.utils.envs_utils import get_pd_split_max_new_tokens
+from lightllm.utils.shm_port_args import get_shm_port_args
from .pd_selector import create_selector
logger = init_logger(__name__)
@@ -34,7 +35,7 @@ def __init__(
self.args = args
self.max_req_total_len = args.max_req_total_len
assert self.max_req_total_len is not None
- self.metric_client = MetricClient(args.metric_port)
+ self.metric_client = MetricClient(get_shm_port_args().metric_port)
self.id_gen = ReqIDGenerator()
self.pd_manager = PDManager(args)
diff --git a/lightllm/server/httpserver_for_pd_master/register_loop.py b/lightllm/server/httpserver_for_pd_master/register_loop.py
index 7e30ffc470..1e1e911b9a 100644
--- a/lightllm/server/httpserver_for_pd_master/register_loop.py
+++ b/lightllm/server/httpserver_for_pd_master/register_loop.py
@@ -5,6 +5,7 @@
from lightllm.utils.net_utils import get_hostname_ip
from lightllm.utils.log_utils import init_logger
from lightllm.server.httpserver_for_pd_master.manager import HttpServerManagerForPDMaster
+from lightllm.utils.shm_port_args import get_shm_port_args
from ..pd_io_struct import PD_Master_Obj
logger = init_logger(__name__)
@@ -18,17 +19,18 @@ async def register_loop(manager: HttpServerManagerForPDMaster):
else:
manager.host_ip = manager.args.host
+ ports = get_shm_port_args()
while True:
try:
- uri = f"ws://{manager.args.config_server_host}:{manager.args.config_server_port}/pd_master_register"
+ uri = f"ws://{manager.args.config_server_host}:{ports.config_server_port}/pd_master_register"
async with websockets.connect(uri, max_queue=(2048 * 1024, 2048 * 1023)) as websocket:
sock = websocket.transport.get_extra_info("socket")
sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1)
pd_master_obj = PD_Master_Obj(
- node_id=manager.args.pd_node_id, host_ip_port=f"{manager.host_ip}:{manager.args.port}"
+ node_id=manager.args.pd_node_id, host_ip_port=f"{manager.host_ip}:{ports.port}"
)
await websocket.send(pickle.dumps(pd_master_obj))
diff --git a/lightllm/server/io_struct.py b/lightllm/server/io_struct.py
new file mode 100644
index 0000000000..011629bca7
--- /dev/null
+++ b/lightllm/server/io_struct.py
@@ -0,0 +1,161 @@
+"""HTTP / Router / Model 之间的 RL 与运维控制面数据结构。
+
+本文件只放「非生成主路径」的类型,结构如下:
+
+1. **Generation control** — ``AbortReq``(多在 HTTP/RL controller 本地处理)
+2. **RL control plane** — 统一经 ``RlOpReq`` 下发到 Model ``RlBackendOps``
+ - **Envelope**:``RlOpReq`` / ``RlOpRsp``(传输信封)
+ - **Op payloads**:具体操作的入参,只是 ``op_name`` 不同
+ - cache / memory:``FlushCacheReq`` / ``ReleaseMemoryReq`` / ``ResumeMemoryReq``
+ - weight update:Init/Destroy group、三种 UpdateWeights*
+
+信封与 payload 不是两类类东西:payload 填进 ``RlOpReq.op_args``,
+由 ``op_name`` 决定后端调用哪个方法。
+
+生成请求的采样/多模态参数不在此文件。
+"""
+
+from dataclasses import dataclass
+from typing import Any, List, Optional, Union
+
+from lightllm.utils.torch_memory_saver_utils import MemoryTag
+
+
+# ---------------------------------------------------------------------------
+# 1. Generation control
+# ---------------------------------------------------------------------------
+
+
+@dataclass
+class AbortReq:
+ """中止指定请求,或中止当前所有请求。
+
+ ``request_id`` 对应内部的 ``group_req_id``。
+ """
+
+ request_id: Optional[int] = None
+ abort_all: bool = False
+
+
+# ---------------------------------------------------------------------------
+# 2. RL control plane
+# ---------------------------------------------------------------------------
+
+
+# 2.1 Envelope — HTTP → Router → Model 的统一传输层
+
+
+@dataclass
+class RlOpReq:
+ """RL 控制面请求信封。
+
+ ``op_name`` 对应 :class:`RlBackendOps` 中的方法名,
+ ``op_args`` 为该方法的入参(下列 Op payload,或已拆好的 tags 等)。
+ """
+
+ op_name: str
+ op_args: Optional[Any] = None
+
+
+@dataclass
+class RlOpRsp:
+ """RL 控制面响应信封。"""
+
+ success: bool
+ msg: Optional[str]
+ op_name: str
+ op_result: Optional[Any] = None
+
+
+# 2.2 Op payloads — cache / memory
+# 与权重更新同属 RlBackendOps 可 dispatch 的操作,仅 op_name / 入参不同。
+
+
+def _normalize_memory_tags(tags):
+ if tags is None:
+ return None
+ return [tag if isinstance(tag, MemoryTag) else MemoryTag(tag) for tag in tags]
+
+
+@dataclass
+class FlushCacheReq:
+ """清空 radix / prompt cache。对应 op: ``flush_cache``。"""
+
+ pass
+
+
+@dataclass
+class ReleaseMemoryReq:
+ """暂停指定 MemoryTag 显存占用。对应 op: ``release_memory_occupation``。"""
+
+ tags: Optional[List[MemoryTag]] = None
+
+ def __post_init__(self):
+ self.tags = _normalize_memory_tags(self.tags)
+
+
+@dataclass
+class ResumeMemoryReq:
+ """恢复先前 pause 的显存占用。对应 op: ``resume_memory_occupation``。"""
+
+ tags: Optional[List[MemoryTag]] = None
+
+ def __post_init__(self):
+ self.tags = _normalize_memory_tags(self.tags)
+
+
+# 2.3 Op payloads — weight update
+
+
+@dataclass
+class InitWeightsUpdateGroupReq:
+ """初始化在线权重更新 process group。对应 op: ``init_weights_update_group``。"""
+
+ master_address: str
+ master_port: int
+ rank_offset: int
+ world_size: int
+ group_name: str = "weight_update_group"
+ backend: str = "nccl"
+
+
+@dataclass
+class DestroyWeightsUpdateGroupReq:
+ """销毁权重更新 process group。对应 op: ``destroy_weights_update_group``。"""
+
+ group_name: str = "weight_update_group"
+
+
+@dataclass
+class UpdateWeightsFromDistributedReq:
+ """经 NCCL group 拉取并更新权重。对应 op: ``update_weights_from_distributed``。"""
+
+ names: List[str]
+ dtypes: List[str]
+ shapes: List[List[int]]
+ group_name: str = "weight_update_group"
+ flush_cache: bool = True
+ abort_all_requests: bool = False
+ weight_version: Optional[str] = None
+
+
+@dataclass
+class UpdateWeightsFromTensorReq:
+ """经序列化 tensor 更新权重。对应 op: ``update_weights_from_tensor``。"""
+
+ serialized_named_tensors: List[Union[str, bytes]]
+ load_format: Optional[str] = None
+ flush_cache: bool = True
+ abort_all_requests: bool = False
+ weight_version: Optional[str] = None
+
+
+@dataclass
+class UpdateWeightsFromIPCReq:
+ """经 CUDA IPC / shm 更新权重。对应 op: ``update_weights_from_ipc``。"""
+
+ ipc_handle: Optional[Union[str, dict]] = None
+ use_shm: bool = False
+ flush_cache: bool = True
+ abort_all_requests: bool = False
+ weight_version: Optional[str] = None
diff --git a/lightllm/server/metrics/manager.py b/lightllm/server/metrics/manager.py
index a95ddc0236..22f6426a77 100644
--- a/lightllm/server/metrics/manager.py
+++ b/lightllm/server/metrics/manager.py
@@ -13,6 +13,7 @@
from lightllm.utils.log_utils import init_logger
from lightllm.utils.graceful_utils import graceful_registry
from lightllm.utils.envs_utils import get_unique_server_name
+from lightllm.utils.shm_port_args import get_shm_port_args
logger = init_logger(__name__)
@@ -158,6 +159,6 @@ def start_metric_manager(args: StartArgs, pipe_writer):
from rpyc.utils.server import ThreadedServer
- t = ThreadedServer(service, port=args.metric_port)
+ t = ThreadedServer(service, port=get_shm_port_args().metric_port)
pipe_writer.send("init ok")
t.start()
diff --git a/lightllm/server/multi_level_kv_cache/manager.py b/lightllm/server/multi_level_kv_cache/manager.py
index 0a7dec0005..ef5b7369c9 100644
--- a/lightllm/server/multi_level_kv_cache/manager.py
+++ b/lightllm/server/multi_level_kv_cache/manager.py
@@ -18,6 +18,7 @@
from lightllm.utils.log_utils import init_logger
from lightllm.utils.process_check import start_parent_check_thread
from lightllm.utils.envs_utils import get_unique_server_name
+from lightllm.utils.shm_port_args import get_shm_port_args
logger = init_logger(__name__)
@@ -28,12 +29,13 @@ def __init__(
args: StartArgs,
):
self.args: StartArgs = args
+ ports = get_shm_port_args()
context = zmq.Context(2)
self.zmq_recv_socket = context.socket(zmq.PULL)
- self.zmq_recv_socket.bind(f"{args.zmq_mode}127.0.0.1:{args.multi_level_kv_cache_port}")
+ self.zmq_recv_socket.bind(f"{args.zmq_mode}127.0.0.1:{ports.multi_level_kv_cache_port}")
self.send_to_router = context.socket(zmq.PUSH)
- self.send_to_router.connect(f"{args.zmq_mode}127.0.0.1:{args.router_port}")
+ self.send_to_router.connect(f"{args.zmq_mode}127.0.0.1:{ports.router_port}")
logger.info(f"send_to_router sendhwm {self.send_to_router.getsockopt(zmq.SNDHWM)}")
self.cpu_cache_client = CpuKvCacheClient(only_create_meta_data=False, init_shm_data=True)
self.shm_req_manager = ShmReqManager()
@@ -162,6 +164,9 @@ def _handle_group_req_multi_cache_match(self, group_req_indexes: GroupReqIndexes
continue
req: Req = req
+ if req.sample_params.prompt_logprobs >= 0:
+ continue
+
token_hash_list = req.token_hash_list.get_all()
if len(token_hash_list) == 0:
continue
diff --git a/lightllm/server/multimodal_params.py b/lightllm/server/multimodal_params.py
index 9541e434c8..68aced6a09 100644
--- a/lightllm/server/multimodal_params.py
+++ b/lightllm/server/multimodal_params.py
@@ -133,7 +133,9 @@ def __init__(self, **kwargs):
async def preload(self, request: Request):
- max_image_pixels = get_env_start_args().max_image_pixels
+ start_args = get_env_start_args()
+ max_image_pixels = start_args.max_image_pixels
+ disable_image_resize = start_args.disable_image_resize
try:
if self._type == "url":
@@ -147,12 +149,15 @@ async def preload(self, request: Request):
# 的 token 计数判断, 所以只需要图片长宽信息,不需要具体图片的内容信息
src_w = self._data[0]
src_h = self._data[1]
- self.image_w, self.image_h = _resize_image_dimensions_if_needed(src_w, src_h, max_image_pixels)
- if (self.image_w, self.image_h) != (src_w, src_h):
- logger.warning(
- f"image_size pixels {src_w * src_h} exceed max_image_pixels={max_image_pixels}, "
- f"resized to {self.image_w}x{self.image_h}"
- )
+ if disable_image_resize:
+ self.image_w, self.image_h = src_w, src_h
+ else:
+ self.image_w, self.image_h = _resize_image_dimensions_if_needed(src_w, src_h, max_image_pixels)
+ if (self.image_w, self.image_h) != (src_w, src_h):
+ logger.warning(
+ f"image_size pixels {src_w * src_h} exceed max_image_pixels={max_image_pixels}, "
+ f"resized to {self.image_w}x{self.image_h}"
+ )
return
else:
raise ValueError(f"cannot read image which type is {self._type}!")
@@ -163,22 +168,24 @@ async def preload(self, request: Request):
loop = asyncio.get_running_loop()
# 1) Verify original input bytes first.
src_w, src_h = await loop.run_in_executor(_IMAGE_VERIFY_POOL, _verify_image_bytes, img_data)
- # 2) Resize (or no-op) after verification.
- img_data, resized_w, resized_h = await loop.run_in_executor(
- _IMAGE_VERIFY_POOL,
- _resize_image_bytes_if_needed,
- img_data,
- src_w,
- src_h,
- max_image_pixels,
- )
- self.image_w, self.image_h = resized_w, resized_h
-
- if (resized_w, resized_h) != (src_w, src_h):
- logger.warning(
- f"image pixels {src_w * src_h} exceed max_image_pixels={max_image_pixels},"
- f" resized to {self.image_w}x{self.image_h}"
+ # 2) Resize after verification unless --disable_image_resize is set.
+ if disable_image_resize:
+ self.image_w, self.image_h = src_w, src_h
+ else:
+ img_data, resized_w, resized_h = await loop.run_in_executor(
+ _IMAGE_VERIFY_POOL,
+ _resize_image_bytes_if_needed,
+ img_data,
+ src_w,
+ src_h,
+ max_image_pixels,
)
+ self.image_w, self.image_h = resized_w, resized_h
+ if (resized_w, resized_h) != (src_w, src_h):
+ logger.warning(
+ f"image pixels {src_w * src_h} exceed max_image_pixels={max_image_pixels},"
+ f" resized to {self.image_w}x{self.image_h}"
+ )
self._preload_data = img_data
return
@@ -236,6 +243,25 @@ async def verify_and_preload(self, request: Request):
await asyncio.gather(*tasks)
return
+ def verify_resource_limits(self):
+ """校验多模态资源数量与 vision/audio 开关是否满足启动配置。
+
+ - images + audios 总数不能超过 ``--cache_capacity``
+ - 若携带 image,则不能处于 ``--disable_vision``
+ - 若携带 audio,则不能处于 ``--disable_audio``
+ """
+ args = get_env_start_args()
+ if len(self.images) + len(self.audios) > args.cache_capacity:
+ raise ValueError(
+ f"too many multimodal items: {len(self.images) + len(self.audios)}"
+ f" > cache_capacity={args.cache_capacity}"
+ )
+ if self.images and args.disable_vision:
+ raise ValueError("vision multimodal not enabled")
+ if self.audios and args.disable_audio:
+ raise ValueError("audio multimodal not enabled")
+ return
+
def to_dict(self):
ret = {}
ret["images"] = [i.to_dict() for i in self.images]
diff --git a/lightllm/server/req_id_generator.py b/lightllm/server/req_id_generator.py
index f7c099c292..8b8d4a5dc5 100644
--- a/lightllm/server/req_id_generator.py
+++ b/lightllm/server/req_id_generator.py
@@ -34,6 +34,9 @@ def __init__(self):
logger.info("ReqIDGenerator init finished")
def _wait_all_workers_ready(self):
+ if self.args.httpserver_workers == 1:
+ return
+
from lightllm.utils.envs_utils import get_unique_server_name
from lightllm.server.core.objs.shm_array import ShmArray
@@ -78,7 +81,11 @@ def _check_and_set_new_id_range(self):
else:
while True:
try:
- config_server_ip_port = f"{self.args.config_server_host}:{self.args.config_server_port}"
+ from lightllm.utils.shm_port_args import get_shm_port_args
+
+ config_server_ip_port = (
+ f"{self.args.config_server_host}:{get_shm_port_args().config_server_port}"
+ )
url = f"http://{config_server_ip_port}/allocate_global_unique_id_range"
response = requests.get(url)
if response.status_code == 200:
diff --git a/lightllm/server/router/dynamic_prompt/linear_att_radix_cache.py b/lightllm/server/router/dynamic_prompt/linear_att_radix_cache.py
index bf07e121e6..6a8e0a3917 100644
--- a/lightllm/server/router/dynamic_prompt/linear_att_radix_cache.py
+++ b/lightllm/server/router/dynamic_prompt/linear_att_radix_cache.py
@@ -470,6 +470,10 @@ def clear_tree_nodes(self):
self.free_radix_cache_to_get_enough_token(need_token_num=self.total_token_num)
return
+ def flush_cache(self):
+ self.free_radix_cache_to_get_enough_token(need_token_num=self.total_token_num)
+ return
+
def deref_to_first_big_page_node(self, node: LinearAttPagedTreeNode) -> Optional[LinearAttPagedTreeNode]:
assert not node.is_big_page_node()
iter_node = node
diff --git a/lightllm/server/router/dynamic_prompt/radix_cache.py b/lightllm/server/router/dynamic_prompt/radix_cache.py
index 21e26c5854..c103a61473 100644
--- a/lightllm/server/router/dynamic_prompt/radix_cache.py
+++ b/lightllm/server/router/dynamic_prompt/radix_cache.py
@@ -106,6 +106,7 @@ class RadixCache:
def __init__(self, unique_name, total_token_num, rank_in_node, mem_manager=None):
from lightllm.common.kv_cache_mem_manager import MemoryManager
+ self.total_token_num = total_token_num
self.mem_manager: MemoryManager = mem_manager
self._key_dtype = torch.int64
self._value_dtype = torch.int64
@@ -419,6 +420,10 @@ def clear_tree_nodes(self):
self.refed_tokens_num.arr[0] = 0
return
+ def flush_cache(self):
+ self.free_radix_cache_to_get_enough_token(need_token_num=self.total_token_num)
+ return
+
def dec_node_ref_counter(self, node: TreeNode):
if node is None:
return
diff --git a/lightllm/server/router/manager.py b/lightllm/server/router/manager.py
index c0560419ff..01634d962e 100644
--- a/lightllm/server/router/manager.py
+++ b/lightllm/server/router/manager.py
@@ -1,7 +1,6 @@
import time
import uvloop
import asyncio
-import torch
import pickle
import inspect
import setproctitle
@@ -12,7 +11,7 @@
import torch.multiprocessing as mp
import torch.distributed as dist
import multiprocessing
-from typing import Dict, List, Optional
+from typing import List
from .batch import Batch, Req
from .model_infer.model_rpc import start_model_process, ModelRpcClient
from .req_queue import build_req_queue
@@ -33,14 +32,17 @@
from lightllm.utils.graceful_utils import graceful_registry
from lightllm.utils.process_check import start_parent_check_thread
from lightllm.utils.envs_utils import get_unique_server_name
+from lightllm.utils.shm_port_args import get_shm_port_args
from lightllm.server.router.dynamic_prompt.shared_arr import SharedInt
from .stats import RouterStatics
from .profiler_service import RouterProfilerCmdQueue, start_router_profiler_server
+from .rl_rpyc import RouterRlOpHelper, start_router_rl_rpyc_server
+from .multinode_tp_helper import RouterMultiNodeTpHelper
logger = init_logger(__name__)
-class RouterManager:
+class RouterManager(RouterMultiNodeTpHelper, RouterRlOpHelper, object):
def __init__(self, args: StartArgs):
self.args = args
self.model_weightdir = args.model_dir
@@ -78,22 +80,23 @@ def __init__(self, args: StartArgs):
self.shared_token_load.set_dynamic_max_load(0.0, dp_index)
self.running_batch: Batch = None
+ ports = get_shm_port_args()
context = zmq.Context(2)
self.zmq_recv_socket = context.socket(zmq.PULL)
- self.zmq_recv_socket.bind(f"{args.zmq_mode}127.0.0.1:{args.router_port}")
+ self.zmq_recv_socket.bind(f"{args.zmq_mode}127.0.0.1:{ports.router_port}")
self.send_to_detokenization = context.socket(zmq.PUSH)
- self.send_to_detokenization.connect(f"{args.zmq_mode}127.0.0.1:{args.detokenization_port}")
+ self.send_to_detokenization.connect(f"{args.zmq_mode}127.0.0.1:{ports.detokenization_port}")
if self.is_multinode_tp:
self.mulitnode_group = dist.init_process_group(
backend="gloo",
- init_method=f"tcp://{args.nccl_host}:{args.multinode_router_gloo_port}",
+ init_method=f"tcp://{args.nccl_host}:{ports.multinode_router_gloo_port}",
world_size=args.nnodes,
rank=args.node_rank,
)
- self.metric_client = MetricClient(args.metric_port)
+ self.metric_client = MetricClient(ports.metric_port)
self.is_pd_run_mode = self.args.run_mode in ["prefill", "decode"]
self.is_pd_decode_mode = self.args.run_mode == "decode"
self.shm_reqs_io_buffer = ShmObjsIOBuffer()
@@ -148,12 +151,11 @@ async def wait_to_model_ready(self):
"max_req_num": self.args.running_max_req_size,
"max_seq_length": self.args.max_req_total_len + 8, # 留一点余量
"nccl_host": self.args.nccl_host,
- "nccl_port": self.args.nccl_port,
+ "nccl_port": get_shm_port_args().nccl_port,
"is_first_token_constraint_mode": self.args.first_token_constraint_mode,
"disable_chunked_prefill": self.args.disable_chunked_prefill,
"chunked_prefill_size": self.args.chunked_prefill_size,
"is_token_healing": self.args.token_healing_mode,
- "return_all_prompt_logprobs": self.args.return_all_prompt_logprobs,
"use_reward_model": self.args.use_reward_model,
"disable_dynamic_prompt_cache": self.args.disable_dynamic_prompt_cache,
"data_type": self.args.data_type,
@@ -167,7 +169,6 @@ async def wait_to_model_ready(self):
"quant_type": self.args.quant_type,
"quant_cfg": self.args.quant_cfg,
"expert_dtype": self.args.expert_dtype,
- "pd_rpyc_ports": self.args.pd_node_infer_rpyc_ports, # 非 pd 模式可以不设置
}
# Call init_model on all model processes
@@ -290,10 +291,17 @@ async def _step(self):
await self._add_batch(new_batch)
self._filter_reqs_from_running_batch()
- aborted_reqs = self._get_aborted_reqs_from_running_batch()
+ # 多机 TP:abort 阶段2(从 running_batch 提取);阶段1 在调度 new_batch 时完成
+ if self.is_multinode_tp:
+ aborted_reqs = self.get_aborted_reqs_from_running_batch_multinode_tp()
+ else:
+ aborted_reqs = self._get_aborted_reqs_from_running_batch()
if aborted_reqs:
await self._aborted_reqs(aborted_reqs=aborted_reqs)
- stop_str_matched_reqs = self._get_stop_str_reqs_from_running_batch()
+ if self.is_multinode_tp:
+ stop_str_matched_reqs = self.get_stop_str_matched_reqs_from_running_batch_multinode_tp()
+ else:
+ stop_str_matched_reqs = self._get_stop_str_reqs_from_running_batch()
if stop_str_matched_reqs:
await self._stop_str_matched_reqs(stop_str_matched_reqs=stop_str_matched_reqs)
return
@@ -350,6 +358,7 @@ def _filter_reqs_from_running_batch(self):
return
def _get_aborted_reqs_from_running_batch(self) -> List[Req]:
+ """非多机 TP:直接读本地 shm 的 is_aborted。"""
ans = []
if self.running_batch is None:
return ans
@@ -360,10 +369,6 @@ def _get_aborted_reqs_from_running_batch(self) -> List[Req]:
return ans
def _get_stop_str_reqs_from_running_batch(self) -> List[Req]:
- # to do, 多节点tp模式,暂时不能支持 stop str 匹配退出
- if self.is_multinode_tp:
- return []
-
ans = []
if self.running_batch is None:
return ans
@@ -433,76 +438,6 @@ def _generate_new_batch(self):
self.schedule_new_batch = Batch.merge_two_batch(self.schedule_new_batch, new_batch)
return
- def _multinode_tp_generate_new_batch(self):
- try:
- dist.barrier(group=self.mulitnode_group)
-
- # 调度的时候需要考虑当前运行的batch,和调度了但是暂时还没有推理的部分请求。
- if self.is_multinode_tp_master:
- new_batch = self.req_queue.generate_new_batch(
- Batch.merge_two_batch(self.running_batch, self.schedule_new_batch)
- )
- if new_batch is not None:
- req_ids = [req.request_id for req in new_batch.reqs]
- else:
- req_ids = []
- dist.broadcast_object_list([len(req_ids)], src=0, group=self.mulitnode_group)
- if len(req_ids) == 0:
- new_batch = None
- else:
- dist.broadcast_object_list(req_ids, src=0, group=self.mulitnode_group)
- req_id_select_mark = [1 for _ in range(len(req_ids))]
- req_id_select_mark = torch.tensor(req_id_select_mark, dtype=torch.int32, device="cpu")
- dist.all_reduce(req_id_select_mark, op=dist.ReduceOp.MIN, group=self.mulitnode_group)
- back_req_list = []
- for req_id, select in zip(req_ids, req_id_select_mark.numpy()):
- if select == 0:
- req = new_batch.pop_req(req_id)
- back_req_list.append(req)
- self.req_queue.waiting_req_list = back_req_list + self.req_queue.waiting_req_list
- if new_batch.is_clear():
- new_batch = None
- else:
- req_nums = [None]
- dist.broadcast_object_list(req_nums, src=0, group=self.mulitnode_group)
- req_num = req_nums[0]
- if req_num == 0:
- new_batch = None
- else:
- req_ids = [None for _ in range(req_num)]
- dist.broadcast_object_list(req_ids, src=0, group=self.mulitnode_group)
- all_req_id_set = set([req.request_id for req in self.req_queue.waiting_req_list])
- req_id_select_mark = []
- for req_id in req_ids:
- req_id_select_mark.append(1 if req_id in all_req_id_set else 0)
- req_id_select_mark = torch.tensor(req_id_select_mark, dtype=torch.int32, device="cpu")
- dist.all_reduce(req_id_select_mark, op=dist.ReduceOp.MIN, group=self.mulitnode_group)
- select_req_ids = []
- for req_id, select in zip(req_ids, req_id_select_mark.numpy()):
- if select == 1:
- select_req_ids.append(req_id)
-
- select_reqs = []
- for req_id in select_req_ids:
- for req in self.req_queue.waiting_req_list:
- if req.request_id == req_id:
- select_reqs.append(req)
-
- for req in select_reqs:
- self.req_queue.waiting_req_list.remove(req)
- if select_reqs:
- new_batch = Batch(-1, reqs=select_reqs, dp_size_in_node=self.dp_size_in_node)
- else:
- new_batch = None
-
- self.schedule_new_batch = Batch.merge_two_batch(self.schedule_new_batch, new_batch)
-
- dist.barrier(group=self.mulitnode_group)
- except Exception as e:
- logger.exception(str(e))
- raise e
- return
-
async def _recv_new_reqs_and_schedule(self):
if not hasattr(self, "recv_max_count"):
self.recv_max_count = 64
@@ -514,7 +449,7 @@ async def _recv_new_reqs_and_schedule(self):
if isinstance(recv_req, GroupReqIndexes):
self._add_req(recv_req)
else:
- assert False, f"Error Req Inf {recv_req}"
+ raise ValueError(f"Unknown request type: {type(recv_req)}")
# 当队列中存在较多的请求时,将一次接受的数量上调
self.recv_max_count = min(int(self.recv_max_count * 1.3), 256)
@@ -523,8 +458,11 @@ async def _recv_new_reqs_and_schedule(self):
# 当队列已经开始清空的时候,将一次接受的数量下调
self.recv_max_count = 64
+ if self.args.enable_rl:
+ await self.process_rl_ops()
+
if self.is_multinode_tp:
- self._multinode_tp_generate_new_batch()
+ self.multinode_tp_generate_new_batch()
else:
if self._get_paused_req_num() == 0:
self._generate_new_batch()
@@ -548,15 +486,16 @@ def handle_exception(loop, context):
asyncio.set_event_loop(loop)
try:
- router = RouterManager(
- args=args,
- )
+ router = RouterManager(args=args)
loop.run_until_complete(router.wait_to_model_ready())
router.profiler_rpyc_server, router.profiler_rpyc_thread = start_router_profiler_server(
args,
router.profiler_cmd_queue,
)
+ router.rl_rpyc_server, router.rl_rpyc_thread = None, None
+ if args.enable_rl:
+ router.rl_rpyc_server, router.rl_rpyc_thread = start_router_rl_rpyc_server(args, router)
except:
import traceback
import sys
diff --git a/lightllm/server/router/model_infer/infer_batch.py b/lightllm/server/router/model_infer/infer_batch.py
index 469f5a289b..ec6e91ac76 100644
--- a/lightllm/server/router/model_infer/infer_batch.py
+++ b/lightllm/server/router/model_infer/infer_batch.py
@@ -23,6 +23,7 @@
from lightllm.utils.envs_utils import get_env_start_args
from lightllm.server.pd_io_struct import PDDecodeNodeInfo
from lightllm.server.embed_cache.embed_cache_client import CpuEmbedCacheClient
+from lightllm.server.router.model_infer.infer_req_ext import FinalTokenMetadataExt, PromptSelectedLogprobsExt
logger = init_logger(__name__)
@@ -264,21 +265,30 @@ def _save_promptcache_kvbuffer(self):
torch.save(prompt_cache_kv_buffer, f"prompt_cache_rank_{dist.get_rank()}.pt")
@torch.no_grad()
- def _filter(self, finished_request_ids: List[int]):
+ def _filter(self, finished_request_ids: List[int], modify_shm_finish_state: bool = True):
if len(finished_request_ids) == 0:
return
+ should_modify_shm = modify_shm_finish_state and self.backend.is_master_in_dp
+
free_req_index = []
free_token_index = []
for request_id in finished_request_ids:
req: InferReq = self.requests_mapping.pop(request_id)
if self.args.diverse_mode:
req.clear_master_slave_state()
+
+ if should_modify_shm:
+ req.final_token_metadata.dump()
+
self.free_a_req_mem(free_token_index, req)
free_req_index.append(req.req_idx)
# logger.info(f"infer release req id {req.shm_req.request_id}")
- req.shm_req.shm_infer_released = True
+ if should_modify_shm:
+ # 释放前兜底:已正常 finished 则 no-op;否则补 finish token 并标 ABORTED。
+ req.mark_shm_aborted_finished()
+ req.shm_req.shm_infer_released = True
self.shm_req_manager.put_back_req_obj(req.shm_req)
if free_token_index:
@@ -594,6 +604,8 @@ def _init_all_state(self):
self.cur_kv_len = 0
self.cur_output_len = 0
+ self.prompt_selected_logprobs = PromptSelectedLogprobsExt(self)
+ self.final_token_metadata = FinalTokenMetadataExt(self)
g_infer_context.req_manager.req_sampling_params_manager.init_req_sampling_params(self)
@@ -712,6 +724,11 @@ def _linear_match_radix_cache(self):
radix_cache.mem_manager.operator.copy_mem_to_mem(
value_tensor[cur_big_page_tokens:shared_kv_len], tail_mems
)
+ # 尾部 KV 换到新 mem 后,同步拷贝已捕获的 top-k prompt logprobs。
+ self.prompt_selected_logprobs.copy_capture_slots_if_needed(
+ source_indexes=value_tensor[cur_big_page_tokens:shared_kv_len],
+ destination_indexes=tail_mems,
+ )
self.shared_kv_node = share_node # 只是为了保证 copy_small_page_buffer_to_linear_att_state 正确调用
g_infer_context.req_manager.copy_small_page_buffer_to_linear_att_state(
@@ -830,10 +847,12 @@ def get_chuncked_input_token_len_for_linear_att(self):
end = self.linear_att_cache_len
return end
- def set_next_gen_token_id(self, next_token_id: int, logprob: float, output_len: int):
+ def set_next_gen_token_id(self, next_token_id: int, logprob: float, output_len: int, rank: int = -1):
index = self.shm_req.input_len + output_len
self.shm_req.shm_prompt_ids.arr[index - 1] = next_token_id
- self.shm_req.shm_logprobs.arr[index - 1] = logprob
+ # structured dtype 整行赋值比分字段 arr["logprob"][i] / arr["rank"][i] 更快
+ # (少两次 field view 查找;bench 约 196ns vs 327ns/次)
+ self.shm_req.shm_logprobs.arr[index - 1] = (logprob, rank)
return
def update_mtp_accepted_token_num(self, accept_token_num: int):
@@ -843,6 +862,47 @@ def update_mtp_accepted_token_num(self, accept_token_num: int):
def get_last_gen_token(self):
return self.shm_req.shm_prompt_ids.arr[self.shm_req.input_len + self.cur_output_len - 1]
+ def mark_shm_aborted_finished(self):
+ """仅写 shm:abort 释放前保证 finish token / finish_status 可用。
+
+ - 本地 ``finish_status`` 已是正常结束:不再改写成 ABORTED。
+ - 尚未结束:补齐 finish token 并标记 FINISHED_ABORTED。
+ - 已有生成:以最后一个输出 token 为 finish token。
+ - 尚无生成:模拟写入 EOS 占位输出。
+ 可读本地状态;不回写本地状态。仅应在 DP master 写 shm 路径调用。
+ """
+ shm_req = self.shm_req
+
+ # 已由正常路径(stop / eos / length 等)正确结束,保留原 finish_status。
+ if self.finish_status.is_finished():
+ return
+
+ input_len = shm_req.input_len
+ output_len = self.cur_output_len
+ if output_len > 0:
+ finish_token_index = input_len + output_len - 1
+ else:
+ eos_ids = self.args.eos_id
+ if eos_ids is None:
+ eos_id = 0
+ elif isinstance(eos_ids, (list, tuple)):
+ eos_id = eos_ids[0]
+ else:
+ eos_id = int(eos_ids)
+ finish_token_index = input_len
+ output_len = 1
+ shm_req.shm_prompt_ids.arr[finish_token_index] = eos_id
+ shm_req.shm_logprobs.arr["logprob"][finish_token_index] = 0.0
+ shm_req.shm_logprobs.arr["rank"][finish_token_index] = -1
+
+ shm_req.finish_token_index = finish_token_index
+ shm_req.finish_status.set_status(FinishStatus.FINISHED_ABORTED)
+
+ shm_req.shm_cur_output_len = output_len
+ # candetoken_out_len 最后写,避免 detoken 提前读到不完整状态
+ shm_req.candetoken_out_len = output_len
+ return
+
def update_finish_status(self, eos_ids, output_len: int):
if self._stop_sequences_matched(output_len=output_len):
self.finish_status.set_status(FinishStatus.FINISHED_STOP)
@@ -901,9 +961,10 @@ def handle(
self,
next_token_id: int,
next_token_logprob: float,
+ next_token_rank: int,
eos_ids: List[int],
- extra_post_req_handle_func: Optional[Callable[[InferReq, int, float], None]],
is_master_in_dp: bool,
+ extra_post_req_handle_func: Optional[Callable[[InferReq, int, float], None]] = None,
pd_prefill_chunked_handle_func: Optional[Callable[[InferReq, int, float, int], None]] = None,
):
# pd_prefill_chunked_handle_func 主要是为了处理 pd prefill 模式下
@@ -917,7 +978,12 @@ def handle(
req_obj = self.req_obj
shm_req = req_obj.shm_req
finish_status = req_obj.finish_status
- req_obj.set_next_gen_token_id(next_token_id, next_token_logprob, self.output_len)
+ req_obj.set_next_gen_token_id(
+ next_token_id,
+ next_token_logprob,
+ self.output_len,
+ rank=next_token_rank,
+ )
# 这里提前判定的主要作用是:
# 在 mtp mode 下,可以存在同一个 req 对象的多次处理,
diff --git a/lightllm/server/router/model_infer/infer_req_ext.py b/lightllm/server/router/model_infer/infer_req_ext.py
new file mode 100644
index 0000000000..14a13ff66a
--- /dev/null
+++ b/lightllm/server/router/model_infer/infer_req_ext.py
@@ -0,0 +1,221 @@
+"""InferReq 功能扩展包裹对象。
+
+将部分与 InferReq 强相关、但不适合继续堆在 InferReq / InferenceContext
+类体上的逻辑拆出,由 InferReq 在初始化时挂载为成员,调用方通过
+``req.`` 访问。
+
+当前包含:
+- :class:`PromptSelectedLogprobsExt`:prompt logprobs 相关辅助
+ (``topk=0`` 异步落盘 + ``topk>0`` mem slot 拷贝)
+- :class:`FinalTokenMetadataExt`:请求结束时汇总并写出 final token metadata
+"""
+
+from typing import TYPE_CHECKING, List, Optional, Tuple
+
+import numpy as np
+import torch
+
+from lightllm.common.basemodel.logprobs_manager import PromptLogprobsCaptureManager
+from lightllm.common.basemodel.moe_route_info_manager import MoeRouteInfoManager
+from lightllm.server.core.objs.token_metadata import ReqFinalTokenMetadata
+
+if TYPE_CHECKING:
+ from lightllm.server.router.model_infer.infer_batch import InferReq
+
+
+class PromptSelectedLogprobsExt:
+ """InferReq 上的 prompt logprobs 辅助对象。
+
+ 覆盖两条路径中与本请求强绑定的操作:
+
+ - ``topk == 0``::meth:`add_chunk` / :meth:`flush` —— 热路径异步 D2H,
+ 结束时写入 ``shm_logprobs``(HTTP 直接读)。
+ - mem 重映射::meth:`copy_capture_slots_if_needed` —— radix / linear
+ att 匹配把 KV 拷到新 mem slot 时,同步拷贝 CaptureManager 中已有的
+ top-k 数据,保证后续请求复用 radix 时元数据仍完整。
+
+ 为何 ``topk==0`` 需要缓冲而不是 prefill 当场写 shm
+ -----------------------------------------------
+ prefill(含 chunked prefill)热路径上立刻 ``.cpu()`` / synchronize 会
+ 打断 overlap。因此只做非阻塞 D2H;``FinalTokenMetadataExt.dump`` →
+ ``flush`` 时再 sync 写 shm。
+
+ chunk 语义(``topk==0``)
+ ------------------------
+ chunked prefill 多次 :meth:`add_chunk`,每段对应
+ ``[target_start, target_end)``;flush 时按段写回。
+
+ 生命周期
+ --------
+ InferReq 初始化时创建;dump metadata 前 flush;必须在
+ ``shm_infer_released=True`` 之前完成,避免 HTTP 读到未写完数据。
+ """
+
+ # (target_start, target_end, logprobs_cpu, ranks_cpu, copy_done_event)
+ _Chunk = Tuple[int, int, torch.Tensor, torch.Tensor, torch.cuda.Event]
+
+ def __init__(self, req: "InferReq") -> None:
+ self._req = req
+ self._chunks: List[PromptSelectedLogprobsExt._Chunk] = []
+
+ def add_chunk(
+ self,
+ target_start: int,
+ target_end: int,
+ logprobs: torch.Tensor,
+ ranks: torch.Tensor,
+ ) -> None:
+ """登记一段 GPU 上的 selected logprobs,异步拷到 pinned CPU。
+
+ Args:
+ target_start: 写入 ``shm_logprobs`` 的起始 prompt 下标(含)。
+ target_end: 写入 ``shm_logprobs`` 的结束 prompt 下标(不含)。
+ logprobs: shape ``[end - start]``,GPU tensor。
+ ranks: shape ``[end - start]``,GPU tensor(1-based rank)。
+ """
+ logprobs_cpu = torch.empty(logprobs.shape, dtype=logprobs.dtype, device="cpu", pin_memory=True)
+ ranks_cpu = torch.empty(ranks.shape, dtype=ranks.dtype, device="cpu", pin_memory=True)
+ logprobs_cpu.copy_(logprobs, non_blocking=True)
+ ranks_cpu.copy_(ranks, non_blocking=True)
+ event = torch.cuda.Event()
+ event.record(torch.cuda.current_stream())
+ self._chunks.append((target_start, target_end, logprobs_cpu, ranks_cpu, event))
+
+ def flush(self) -> None:
+ """等待所有异步 D2H 完成,将结果写入 ``shm_logprobs``。
+
+ 调用时机:请求结束、写出 final token metadata 之前。
+ HTTP 进程对 ``prompt_logprobs=0`` 依赖 ``shm_logprobs``,因此必须在
+ 标记 ``shm_infer_released`` 之前完成提交。
+ """
+ if not self._chunks:
+ return
+
+ shm_logprobs = self._req.shm_req.shm_logprobs
+ for target_start, target_end, logprobs_cpu, ranks_cpu, event in self._chunks:
+ event.synchronize()
+ shm_logprobs.arr["logprob"][target_start:target_end] = logprobs_cpu.numpy()
+ shm_logprobs.arr["rank"][target_start:target_end] = ranks_cpu.numpy()
+
+ self._chunks.clear()
+
+ def copy_capture_slots_if_needed(
+ self,
+ source_indexes: torch.Tensor,
+ destination_indexes: torch.Tensor,
+ ) -> None:
+ """KV mem slot 重映射时,拷贝 CaptureManager 中的 prompt top-k 数据。
+
+ 场景:linear att radix 小页命中后,尾部 KV 从共享小页 mem 拷到新申请的
+ ``tail_mems``。top-k prompt logprobs 按 mem slot 索引存放,必须随 KV
+ 一并拷到新 slot。
+
+ 不能按「当前请求的 prompt_logprobs」决定是否拷贝:source slot 上的
+ 数据可能来自更早的 ``topk>0`` 请求;destination 之后还可能插回
+ radix cache,被后续 ``topk>0`` 请求复用。若此处因当前 ``topk<=0``
+ 跳过拷贝,后续复用方会 extract 到空/错误数据。只要 capture buffer
+ 已初始化,就按 ``max_topk`` 全宽拷贝以保留可复用元数据。
+ """
+ mgr = PromptLogprobsCaptureManager.get_instance()
+ if mgr is None or not mgr.is_buffer_initialized():
+ return
+
+ mgr.copy_slots(
+ source_indexes=source_indexes,
+ destination_indexes=destination_indexes,
+ topk=mgr.max_topk,
+ )
+
+
+class FinalTokenMetadataExt:
+ """请求结束时汇总可选元信息,并写入 final token metadata shm。
+
+ 调用时机
+ --------
+ 仅 DP master 在 ``InferenceContext._filter(modify_shm_finish_state=True)``
+ 路径、真正释放请求前调用 :meth:`dump`。必须在 ``shm_infer_released=True``
+ 之前完成,供 HTTP 侧编码进响应。
+
+ 与 :class:`PromptSelectedLogprobsExt` 的分工
+ ------------------------------------------
+ - ``prompt_logprobs=0``:热路径写入 ``PromptSelectedLogprobsExt``,dump
+ 时 flush 到 ``shm_logprobs``(HTTP 直接读该 shm)。
+ - ``prompt_logprobs>0`` / routed experts:prefill/decode 期间按 mem slot
+ 落在对应 CaptureManager;dump 时按本请求 mem_indexes extract,再与
+ 其它字段一并 pickle 进 metadata shm。
+ """
+
+ def __init__(self, req: "InferReq") -> None:
+ self._req = req
+
+ def _mem_indexes(self) -> torch.Tensor:
+ # 延迟导入,避免与 infer_batch 形成模块级循环依赖。
+ from lightllm.server.router.model_infer.infer_batch import g_infer_context
+
+ return g_infer_context.req_manager.req_to_token_indexs[self._req.req_idx]
+
+ def collect_prompt_logprobs(self) -> Optional[Tuple[np.ndarray, np.ndarray]]:
+ """从 PromptLogprobsCaptureManager 提取 ``prompt_logprobs>0`` 的 top-k。
+
+ Returns:
+ ``(top_token_ids, top_logprobs)``,或无需收集时返回 ``None``
+ (``topk<=0`` / prompt 过短)。
+ """
+ req = self._req
+ topk = req.sampling_param.shm_param.prompt_logprobs
+ if topk <= 0 or req.shm_req.input_len <= 1:
+ return None
+
+ # prefill 未跑完(含中途 abort):不满 input_len-1,不导出 logprobs
+ if req.cur_kv_len < req.shm_req.input_len:
+ return None
+
+ mgr = PromptLogprobsCaptureManager.get_instance()
+ mem_indexes = self._mem_indexes()[: req.shm_req.input_len - 1]
+ return mgr.extract(mem_indexes, topk)
+
+ def collect_routed_experts(self) -> Optional[np.ndarray]:
+ """从 MoeRouteInfoManager 提取已完成请求的 routed expert 信息。"""
+ req = self._req
+
+ visible_total_len = req.shm_req.input_len + req.cur_output_len
+ capture_len = min(req.cur_kv_len, visible_total_len - 1)
+ if capture_len <= 0:
+ return None
+
+ mem_indexes = self._mem_indexes()[0:capture_len]
+ return MoeRouteInfoManager.get_instance().extract(mem_indexes)
+
+ def dump(self) -> None:
+ """汇总各路径元信息并写入 final token metadata shm。"""
+ req = self._req
+ prompt_top_token_ids = None
+ prompt_top_logprobs = None
+ routed_experts = None
+
+ # 阶段 1:落盘 prompt_logprobs=0 的异步缓冲。
+ # prefill 热路径只做了非阻塞 D2H,这里 sync 后写入 shm_logprobs,
+ # HTTP 对该模式直接从 shm_logprobs 读 logprob/rank。
+ req.prompt_selected_logprobs.flush()
+
+ # 阶段 2:收集 prompt_logprobs>0 的 top-k 结果。
+ # 数据按 KV mem slot 存在 PromptLogprobsCaptureManager 中,
+ # 按当前 req 的 mem_indexes extract 后交给 metadata shm。
+ prompt_logprobs_mgr = PromptLogprobsCaptureManager.get_instance()
+ if prompt_logprobs_mgr is not None and prompt_logprobs_mgr.is_buffer_initialized():
+ collected = self.collect_prompt_logprobs()
+ if collected is not None:
+ prompt_top_token_ids, prompt_top_logprobs = collected
+
+ # 阶段 3:收集 MoE routed experts(若开启 enable_return_routed_experts)。
+ # 同样按 mem slot 从 MoeRouteInfoManager 导出。
+ moe_mgr = MoeRouteInfoManager.get_instance()
+ if moe_mgr is not None and moe_mgr.is_buffer_initialized():
+ routed_experts = self.collect_routed_experts()
+
+ # 阶段 4:统一 pickle 写入 final token metadata shm,供 HTTP 编码进响应。
+ ReqFinalTokenMetadata(req.shm_req).save(
+ prompt_top_token_ids=prompt_top_token_ids,
+ prompt_top_logprobs=prompt_top_logprobs,
+ routed_experts=routed_experts,
+ )
diff --git a/lightllm/server/router/model_infer/mode_backend/__init__.py b/lightllm/server/router/model_infer/mode_backend/__init__.py
index 1a4bf6c020..8c608e5f92 100644
--- a/lightllm/server/router/model_infer/mode_backend/__init__.py
+++ b/lightllm/server/router/model_infer/mode_backend/__init__.py
@@ -1,7 +1,6 @@
from .chunked_prefill.impl import ChunkedPrefillBackend
from .chunked_prefill.impl_for_first_token_constraint_mode import FirstTokenConstraintBackend
from .chunked_prefill.impl_for_outlines_constraint_mode import OutlinesConstraintBackend
-from .chunked_prefill.impl_for_return_all_prompt_logprobs import ReturnPromptLogProbBackend
from .chunked_prefill.impl_for_reward_model import RewardModelBackend
from .chunked_prefill.impl_for_token_healing import TokenHealingBackend
from .chunked_prefill.impl_for_xgrammar_mode import XgrammarBackend
diff --git a/lightllm/server/router/model_infer/mode_backend/base_backend.py b/lightllm/server/router/model_infer/mode_backend/base_backend.py
index 3858cdb0b3..eaf3552607 100644
--- a/lightllm/server/router/model_infer/mode_backend/base_backend.py
+++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py
@@ -4,7 +4,7 @@
import time
import threading
import torch.distributed as dist
-from typing import List, Tuple, Callable, Optional
+from typing import List, Tuple, Callable, Optional, Union
from transformers.configuration_utils import PretrainedConfig
from lightllm.utils.infer_utils import set_random_seed
from lightllm.utils.log_utils import init_logger
@@ -12,11 +12,13 @@
from lightllm.server.router.model_infer.infer_batch import InferReq, InferReqUpdatePack
from lightllm.server.router.token_load import TokenLoad
from lightllm.common.basemodel.basemodel import TpPartBaseModel
+from lightllm.common.basemodel.logprobs_manager import PromptLogprobsCaptureManager
+from lightllm.common.basemodel.moe_route_info_manager import MoeRouteInfoManager
from lightllm.common.req_manager import ReqManagerForMamba
from lightllm.common.linear_att_cache_manager import LinearAttCacheManager
from lightllm.server.router.dynamic_prompt.linear_att_radix_cache import LinearAttPagedRadixCache
from lightllm.server.router.dynamic_prompt.radix_cache import RadixCache
-from lightllm.common.basemodel.batch_objs import ModelOutput
+from lightllm.common.basemodel.batch_objs import ModelOutput, ModelInput
from lightllm.common.basemodel.triton_kernel.mtp_utils import mtp_verify
from lightllm.utils.dist_utils import init_distributed_env
from lightllm.utils.envs_utils import get_unique_server_name
@@ -33,8 +35,8 @@
enable_radix_tree_timer_merge,
get_radix_tree_merge_update_delta,
)
+from lightllm.distributed import dist_group_manager
from lightllm.distributed.communication_op import (
- dist_group_manager,
all_gather_into_tensor,
all_reduce,
broadcast,
@@ -99,7 +101,6 @@ def init_model(self, kvargs):
self.load_way = kvargs["load_way"]
self.disable_chunked_prefill = self.args.disable_chunked_prefill
self.chunked_prefill_size = self.args.chunked_prefill_size
- self.return_all_prompt_logprobs = self.args.return_all_prompt_logprobs
self.use_dynamic_prompt_cache = not self.args.disable_dynamic_prompt_cache
self.batch_max_tokens = self.args.batch_max_tokens
self.eos_id: List[int] = kvargs.get("eos_id", [2])
@@ -111,8 +112,6 @@ def init_model(self, kvargs):
self.logger = init_logger(__name__)
self.weight_dir = kvargs["weight_dir"]
- # p d 分离模式,decode节点才会使用的参数
- self.pd_rpyc_ports = kvargs.get("pd_rpyc_ports", None)
max_total_token_num = kvargs["max_total_token_num"]
init_distributed_env(kvargs)
@@ -136,7 +135,7 @@ def init_model(self, kvargs):
"max_req_num": kvargs.get("max_req_num", 1000),
"max_seq_length": kvargs.get("max_seq_length", 1024 * 5),
"is_token_healing": kvargs.get("is_token_healing", False),
- "return_all_prompt_logics": self.return_all_prompt_logprobs,
+ "return_all_prompt_logics": self.args.enable_prompt_logprobs,
"disable_chunked_prefill": self.disable_chunked_prefill,
"data_type": kvargs.get("data_type", "float16"),
"graph_max_batch_size": kvargs.get("graph_max_batch_size", 16),
@@ -196,7 +195,6 @@ def init_model(self, kvargs):
shm_req_manager=self.shm_req_manager,
vocab_size=self.model.vocab_size,
)
-
# 初始化 dp 模式使用的通信 tensor, 对于非dp模式,不会使用到
if self.dp_size > 1:
self.dp_reduce_tensor = torch.tensor([0], dtype=torch.int32, device="cuda", requires_grad=False)
@@ -225,6 +223,19 @@ def init_model(self, kvargs):
self.model.mem_manager.write_to_shm(req_manager=self.model.req_manager)
dist.barrier(group=self.node_nccl_group)
+ # 同一 DP 组内只需主 rank 初始化真实的 capture buffer 并执行后续相关操作;
+ # 非主 rank 不需要分配 buffer,避免重复占用内存。
+ if self.is_master_in_dp:
+ kv_cache_size = self.model.mem_manager.size + 1
+ if self.args.enable_prompt_logprobs:
+ mgr = PromptLogprobsCaptureManager.get_instance()
+ if mgr is not None:
+ mgr.init_capture_buffer(kv_cache_size=kv_cache_size)
+ if self.args.enable_return_routed_experts:
+ mgr = MoeRouteInfoManager.get_instance()
+ if mgr is not None:
+ mgr.init_capture_buffer(kv_cache_size=kv_cache_size)
+
self.init_custom()
if self.args.enable_dp_prompt_cache_fetch:
@@ -357,10 +368,15 @@ def init_mtp_draft_model(self, main_kvargs: dict):
self.logger.info(f"loaded mtp model class {self.draft_models[i].__class__}")
return
- def _async_copy_next_token_infos_to_pin_mem(self, next_token_ids: torch.Tensor, next_token_logprobs: torch.Tensor):
+ def _async_copy_next_token_infos_to_pin_mem(
+ self,
+ next_token_ids: torch.Tensor,
+ next_token_logprobs: torch.Tensor,
+ next_token_ranks: torch.Tensor,
+ ):
"""
- 这个函数会把next token id和logprobs保存到pinned memory中
- 这样可以保障post_handle 函数可以读取到正常的输出结果。
+ 把 next token id / logprobs / ranks 异步拷到 pinned memory,
+ 供后续 post_handle 读取。ranks 始终有值(不需要时为常量 -1)。
"""
next_token_ids_cpu = g_pin_mem_manager.async_copy_from_gpu_tensor(
key="next_token_ids",
@@ -370,7 +386,89 @@ def _async_copy_next_token_infos_to_pin_mem(self, next_token_ids: torch.Tensor,
key="next_token_logprobs",
gpu_tensor=next_token_logprobs,
)
- return next_token_ids_cpu, next_token_logprobs_cpu
+ # 仅 enable_rl 需要真实 rank;否则跳过 D2H,返回常量 -1。
+ if self.args.enable_rl:
+ next_token_ranks_cpu = g_pin_mem_manager.async_copy_from_gpu_tensor(
+ key="next_token_ranks",
+ gpu_tensor=next_token_ranks,
+ )
+ else:
+ next_token_ranks_cpu = g_pin_mem_manager.get_const_cpu_tensor(
+ key="next_token_ranks",
+ shape=next_token_ids_cpu.shape,
+ fill_value=-1,
+ dtype=torch.int32,
+ )
+ return next_token_ids_cpu, next_token_logprobs_cpu, next_token_ranks_cpu
+
+ def _get_next_token_ranks(self, logits: torch.Tensor, next_token_ids: torch.Tensor) -> torch.Tensor:
+ """计算(或占位)每个 next token 在 vocab 上的 1-based rank(GPU tensor)。
+
+ 仅 ``--enable_rl`` 时做真实 rank;否则返回 GPU 常量 ``-1``,避免 O(batch * vocab) 比较。
+ 下游 async_copy 在同样条件下会忽略该返回值。
+ """
+ if not self.args.enable_rl:
+ return g_pin_mem_manager.get_const_gpu_tensor(
+ key="next_token_ranks",
+ shape=next_token_ids.shape,
+ fill_value=-1,
+ dtype=torch.int32,
+ )
+ selected_logits = logits.gather(1, next_token_ids.long().view(-1, 1))
+ return (logits > selected_logits).sum(dim=-1, dtype=torch.int32) + 1
+
+ def _capture_prompt_logprobs_if_needed(
+ self,
+ model_input: ModelInput,
+ run_reqs: List[InferReq],
+ prompt_logits: Optional[torch.Tensor],
+ ) -> None:
+ # 仅在开启 return_all_prompt_logics(如 enable_prompt_logprobs)且存在完整的
+ # prefill logits 时,才需要捕获每个 prompt token 的 logprobs 信息。此时用于采样
+ # 的 logits 已经只对应每个请求最后一个位置,无需再处理。
+ if not self.model.return_all_prompt_logics or prompt_logits is None:
+ return
+
+ mgr = PromptLogprobsCaptureManager.get_instance()
+
+ start_loc = 0
+ for req_obj in run_reqs:
+ q_len = req_obj.prefill_need_token_num(is_chuncked_prefill=not self.disable_chunked_prefill)
+ topk = req_obj.sampling_param.shm_param.prompt_logprobs
+ capture_count = min(q_len, req_obj.shm_req.input_len - req_obj.cur_kv_len - 1)
+ if capture_count > 0 and topk == 0 and self.is_master_in_dp:
+ # prompt_logprobs=0 返回真实命中的 prompt token,
+ # rank 必须基于全 vocab 计算,不能用 top-k 列表位置替代。
+ logit_rows = prompt_logits[start_loc : start_loc + capture_count]
+ target_start = req_obj.cur_kv_len + 1
+ target_end = target_start + capture_count
+ target_token_ids = torch.tensor(
+ req_obj.shm_req.shm_prompt_ids.arr[target_start:target_end].copy(),
+ dtype=torch.long,
+ device=logit_rows.device,
+ )
+ target_logits = logit_rows.gather(1, target_token_ids.long().view(-1, 1)).view(-1)
+ logprobs = target_logits.float() - torch.logsumexp(logit_rows.float(), dim=-1)
+ ranks = (logit_rows > target_logits.view(-1, 1)).sum(dim=-1, dtype=torch.int32) + 1
+ req_obj.prompt_selected_logprobs.add_chunk(target_start, target_end, logprobs, ranks)
+ elif capture_count > 0 and topk > 0 and mgr is not None and mgr.is_buffer_initialized():
+ logit_rows = prompt_logits[start_loc : start_loc + capture_count]
+ log_normalizer = torch.logsumexp(logit_rows.float(), dim=-1)
+ valid_topk = min(topk, logit_rows.shape[-1])
+ top_logits, top_token_ids = logit_rows.topk(valid_topk, dim=-1)
+ top_token_ids = top_token_ids.to(torch.int32)
+ top_logprobs = top_logits.float() - log_normalizer.view(-1, 1)
+ if valid_topk < topk:
+ padding = (0, topk - valid_topk)
+ top_token_ids = torch.nn.functional.pad(top_token_ids, padding, value=-1)
+ top_logprobs = torch.nn.functional.pad(top_logprobs, padding, value=float("-inf"))
+ mgr.capture(
+ mem_indexes=model_input.mem_indexes[start_loc : start_loc + capture_count],
+ top_token_ids=top_token_ids,
+ top_logprobs=top_logprobs,
+ )
+ start_loc += q_len
+ return
def _try_read_new_reqs(self):
if self.is_multinode_tp:
@@ -481,10 +579,9 @@ def _read_pd_trans_io_buffer_and_update_req_status(self):
InferReqUpdatePack(req_obj=req, output_len=req.cur_output_len).handle(
next_token_id=obj.first_gen_token_id,
next_token_logprob=obj.first_gen_token_logprob,
+ next_token_rank=-1,
eos_ids=self.eos_id,
- extra_post_req_handle_func=None,
is_master_in_dp=self.is_master_in_dp,
- pd_prefill_chunked_handle_func=None,
)
return
@@ -722,6 +819,7 @@ def _post_handle(
run_reqs: List[InferReq],
next_token_ids: List[int],
next_token_logprobs: List[float],
+ next_token_ranks: List[int],
run_reqs_update_packs: List[InferReqUpdatePack],
extra_post_req_handle_func: Optional[Callable[[InferReq, int, float], None]] = None,
pd_prefill_chunked_handle_func: Optional[Callable[[InferReq, int, float, int], None]] = None,
@@ -730,17 +828,18 @@ def _post_handle(
extra_post_req_handle_func 用于提供在一个请求确定输出的时候,给出额外的后处理操作,主要是用于
约束输出等模式,设置自己请求内部的状态机的状态,并添加额外的停止判定条件等。
"""
- for req_obj, next_token_id, next_token_logprob, pack in zip(
- run_reqs, next_token_ids, next_token_logprobs, run_reqs_update_packs
+ for req_obj, next_token_id, next_token_logprob, next_token_rank, pack in zip(
+ run_reqs, next_token_ids, next_token_logprobs, next_token_ranks, run_reqs_update_packs
):
req_obj: InferReq = req_obj
pack: InferReqUpdatePack = pack
pack.handle(
next_token_id=next_token_id,
next_token_logprob=next_token_logprob,
+ next_token_rank=int(next_token_rank),
eos_ids=self.eos_id,
- extra_post_req_handle_func=extra_post_req_handle_func,
is_master_in_dp=self.is_master_in_dp,
+ extra_post_req_handle_func=extra_post_req_handle_func,
pd_prefill_chunked_handle_func=pd_prefill_chunked_handle_func,
)
@@ -801,6 +900,7 @@ def _sample_and_scatter_token(
mask_func(run_reqs, logits)
next_token_ids, next_token_logprobs = sample(logits, run_reqs, self.eos_id)
+ next_token_ranks = self._get_next_token_ranks(logits, next_token_ids)
b_has_out = None
if is_prefill:
b_has_out = g_pin_mem_manager.gen_from_list(
@@ -819,10 +919,16 @@ def _sample_and_scatter_token(
next_token_ids=next_token_ids,
mask=b_has_out,
)
- next_token_ids_cpu, next_token_logprobs_cpu = self._async_copy_next_token_infos_to_pin_mem(
- next_token_ids, next_token_logprobs
+ (
+ next_token_ids_cpu,
+ next_token_logprobs_cpu,
+ next_token_ranks_cpu,
+ ) = self._async_copy_next_token_infos_to_pin_mem(
+ next_token_ids,
+ next_token_logprobs,
+ next_token_ranks,
)
- return next_token_ids, next_token_ids_cpu, next_token_logprobs_cpu
+ return next_token_ids, next_token_ids_cpu, next_token_logprobs_cpu, next_token_ranks_cpu
def _dp_all_gather_prefill_and_decode_req_num(
self, prefill_reqs: List[InferReq], decode_reqs: List[InferReq]
diff --git a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py
index 967a4c150f..d75302800b 100644
--- a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py
+++ b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py
@@ -109,7 +109,8 @@ def prefill_normal(
model_input, run_reqs = prepare_prefill_inputs(prefill_reqs, is_chuncked_mode=not self.disable_chunked_prefill)
with torch.cuda.stream(g_infer_context.get_overlap_stream()):
model_output = self.model.forward(model_input)
- _, next_token_ids_cpu, next_token_logprobs_cpu = self._sample_and_scatter_token(
+ self._capture_prompt_logprobs_if_needed(model_input, run_reqs, model_output.prompt_logics)
+ (_, next_token_ids_cpu, next_token_logprobs_cpu, next_token_ranks_cpu,) = self._sample_and_scatter_token(
logits=model_output.logits,
b_req_idx=model_input.b_req_idx,
b_mtp_index=model_input.b_mtp_index,
@@ -136,6 +137,7 @@ def prefill_normal(
run_reqs=run_reqs,
next_token_ids=next_token_ids_cpu,
next_token_logprobs=next_token_logprobs_cpu,
+ next_token_ranks=next_token_ranks_cpu,
run_reqs_update_packs=update_packs,
extra_post_req_handle_func=self.extra_post_req_handle_func,
pd_prefill_chunked_handle_func=self.pd_prefill_chunked_handle_func,
@@ -152,7 +154,7 @@ def decode_normal(
model_input, run_reqs = prepare_decode_inputs(decode_reqs)
with torch.cuda.stream(g_infer_context.get_overlap_stream()):
model_output = self.model.forward(model_input)
- _, next_token_ids_cpu, next_token_logprobs_cpu = self._sample_and_scatter_token(
+ (_, next_token_ids_cpu, next_token_logprobs_cpu, next_token_ranks_cpu,) = self._sample_and_scatter_token(
logits=model_output.logits,
b_req_idx=model_input.b_req_idx,
b_mtp_index=model_input.b_mtp_index,
@@ -174,6 +176,7 @@ def decode_normal(
run_reqs=run_reqs,
next_token_ids=next_token_ids_cpu,
next_token_logprobs=next_token_logprobs_cpu,
+ next_token_ranks=next_token_ranks_cpu,
run_reqs_update_packs=update_packs,
extra_post_req_handle_func=self.extra_post_req_handle_func,
)
@@ -190,7 +193,13 @@ def prefill_mtp(
model_input, run_reqs = prepare_prefill_inputs(prefill_reqs, is_chuncked_mode=not self.disable_chunked_prefill)
with torch.cuda.stream(g_infer_context.get_overlap_stream()):
model_output = self.model.forward(model_input)
- next_token_ids, next_token_ids_cpu, next_token_logprobs_cpu = self._sample_and_scatter_token(
+ self._capture_prompt_logprobs_if_needed(model_input, run_reqs, model_output.prompt_logics)
+ (
+ next_token_ids,
+ next_token_ids_cpu,
+ next_token_logprobs_cpu,
+ next_token_ranks_cpu,
+ ) = self._sample_and_scatter_token(
logits=model_output.logits,
b_req_idx=model_input.b_req_idx,
b_mtp_index=model_input.b_mtp_index,
@@ -222,6 +231,7 @@ def prefill_mtp(
run_reqs=run_reqs,
next_token_ids=next_token_ids_cpu,
next_token_logprobs=next_token_logprobs_cpu,
+ next_token_ranks=next_token_ranks_cpu,
run_reqs_update_packs=update_packs,
extra_post_req_handle_func=self.extra_post_req_handle_func,
pd_prefill_chunked_handle_func=self.pd_prefill_chunked_handle_func,
@@ -245,6 +255,7 @@ def decode_mtp(
b_mtp_index_cpu = model_input.b_mtp_index
model_output = self.model.forward(model_input)
next_token_ids, next_token_logprobs = sample(model_output.logits, run_reqs, self.eos_id)
+ next_token_ranks = self._get_next_token_ranks(model_output.logits, next_token_ids)
# verify the next_token_ids
b_req_mtp_start_loc = [index for index, mtp_index in enumerate(b_mtp_index_cpu) if mtp_index == 0]
b_req_mtp_start_loc = g_pin_mem_manager.gen_from_list(
@@ -278,9 +289,11 @@ def decode_mtp(
verify_event = torch.cuda.Event()
verify_event.record()
- next_token_ids_cpu, next_token_logprobs_cpu = self._async_copy_next_token_infos_to_pin_mem(
- next_token_ids, next_token_logprobs
- )
+ (
+ next_token_ids_cpu,
+ next_token_logprobs_cpu,
+ next_token_ranks_cpu,
+ ) = self._async_copy_next_token_infos_to_pin_mem(next_token_ids, next_token_logprobs, next_token_ranks)
# 调用具体的draft decode函数
additional_mem_indexes_cpu = self._draft_decode_func(
@@ -320,6 +333,7 @@ def decode_mtp(
run_reqs=verify_ok_reqs,
next_token_ids=next_token_ids_cpu[select_mask],
next_token_logprobs=next_token_logprobs_cpu[select_mask],
+ next_token_ranks=next_token_ranks_cpu[select_mask],
run_reqs_update_packs=update_packs,
extra_post_req_handle_func=self.extra_post_req_handle_func,
)
diff --git a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl_for_return_all_prompt_logprobs.py b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl_for_return_all_prompt_logprobs.py
deleted file mode 100644
index a9fbb41e7f..0000000000
--- a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl_for_return_all_prompt_logprobs.py
+++ /dev/null
@@ -1,68 +0,0 @@
-import torch
-from .impl import ChunkedPrefillBackend
-from typing import List
-from lightllm.server.router.model_infer.infer_batch import InferReq
-from lightllm.server.router.model_infer.mode_backend.pre import prepare_prefill_inputs
-from lightllm.server.router.model_infer.mode_backend.generic_post_process import sample
-from lightllm.server.router.model_infer.mode_backend.overlap_events import OverlapEventPack
-
-
-class ReturnPromptLogProbBackend(ChunkedPrefillBackend):
- def __init__(self) -> None:
- super().__init__()
- self.prefill = self.return_all_prompt_logprobs_prefill
- return
-
- def return_all_prompt_logprobs_prefill(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq]):
-
- # 在 return all_prompt_logprobs 的模式下,不能启用 dynamic prompt cache
- assert self.radix_cache is None
- assert self.disable_chunked_prefill is True
-
- model_input, run_reqs = prepare_prefill_inputs(prefill_reqs, is_chuncked_mode=not self.disable_chunked_prefill)
-
- model_output = self.model.forward(model_input)
- prompt_all_logits = model_output.logits
-
- input_ids = model_input.input_ids
- b_ready_cache_len = model_input.b_ready_cache_len
- b_seq_len = model_input.b_seq_len
- last_index = torch.cumsum(b_seq_len, dim=0, dtype=torch.long) - 1
- logits = prompt_all_logits[last_index, :]
-
- b_q_seq_len = b_seq_len - b_ready_cache_len
- b_start_loc = torch.cumsum(b_q_seq_len, dim=0, dtype=torch.long) - b_q_seq_len
- b_start_loc = b_start_loc.cpu().numpy()
- b_q_seq_len = b_q_seq_len.cpu().numpy()
-
- for req_obj, start_loc, q_seq_len in zip(run_reqs, b_start_loc, b_q_seq_len):
- req_obj: InferReq = req_obj
- cur_ids: torch.Tensor = input_ids[start_loc : start_loc + q_seq_len]
- cur_logits = prompt_all_logits[start_loc : start_loc + q_seq_len]
- cur_logprobs = torch.log_softmax(cur_logits, dim=-1, dtype=torch.float)[0:-1, :]
- cur_logprobs = torch.gather(cur_logprobs, dim=1, index=cur_ids[1:].view(-1, 1)).detach().cpu().numpy()
-
- if req_obj.shm_req.input_len > 1:
- if self.is_master_in_dp:
- req_obj.shm_req.shm_logprobs.arr[1 : req_obj.shm_req.input_len] = cur_logprobs.flatten()
-
- if self.prefill_mask_func is not None:
- self.prefill_mask_func(run_reqs, logits)
-
- next_token_ids, next_token_probs = sample(logits, run_reqs, self.eos_id)
- next_token_ids = next_token_ids.detach().cpu().numpy()
- next_token_logprobs = torch.log(next_token_probs).detach().cpu().numpy()
-
- update_packs = self._pre_post_handle(run_reqs, is_chuncked_mode=not self.disable_chunked_prefill)
- self._post_handle(
- run_reqs=run_reqs,
- next_token_ids=next_token_ids,
- next_token_logprobs=next_token_logprobs,
- run_reqs_update_packs=update_packs,
- extra_post_req_handle_func=self.extra_post_req_handle_func,
- )
-
- event_pack.notify_post_handle_and_wait_pre_post_handle()
- event_pack.notify_forward_and_wait_post_handle()
- event_pack.notify_pre_post_handle()
- return
diff --git a/lightllm/server/router/model_infer/mode_backend/diverse_backend/impl.py b/lightllm/server/router/model_infer/mode_backend/diverse_backend/impl.py
index f1681eda52..1edbd30306 100644
--- a/lightllm/server/router/model_infer/mode_backend/diverse_backend/impl.py
+++ b/lightllm/server/router/model_infer/mode_backend/diverse_backend/impl.py
@@ -63,6 +63,7 @@ def beam_prefill(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq
b_mtp_index = model_input.b_mtp_index[batch_idx]
next_token_ids, next_token_logprobs = sample(logits, run_reqs, self.eos_id)
+ next_token_ranks = self._get_next_token_ranks(logits, next_token_ids)
scatter_token(
next_token_ids=next_token_ids,
@@ -72,8 +73,14 @@ def beam_prefill(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq
b_has_out=b_has_out,
)
- next_token_ids_cpu, next_token_logprobs_cpu = self._async_copy_next_token_infos_to_pin_mem(
- next_token_ids=next_token_ids, next_token_logprobs=next_token_logprobs
+ (
+ next_token_ids_cpu,
+ next_token_logprobs_cpu,
+ next_token_ranks_cpu,
+ ) = self._async_copy_next_token_infos_to_pin_mem(
+ next_token_ids=next_token_ids,
+ next_token_logprobs=next_token_logprobs,
+ next_token_ranks=next_token_ranks,
)
sync_event = torch.cuda.Event()
@@ -90,6 +97,7 @@ def beam_prefill(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq
run_reqs=run_reqs,
next_token_ids=next_token_ids_cpu,
next_token_logprobs=next_token_logprobs_cpu,
+ next_token_ranks=next_token_ranks_cpu,
run_reqs_update_packs=update_packs,
extra_post_req_handle_func=self.extra_post_req_handle_func,
)
diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py
index 399f797987..f6ca89e651 100644
--- a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py
+++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py
@@ -78,7 +78,14 @@ def _init_reqs(self, reqs: List[Tuple]):
trans_taskes = self.dp_kv_shared_module.build_shared_kv_trans_tasks(reqs=infer_reqs, req_dp_ranks=req_dp_ranks)
self.dp_kv_shared_module.kv_trans(trans_tasks=trans_taskes)
- g_infer_context._filter(finished_request_ids=[req[0] for req in other_dp_reqs])
+ # other_dp_reqs 只是为本 DP 做完 prefix cache / KV 拉取后的临时本地对象,
+ # 真正推理仍在其归属 DP 上进行。这里仅清理本 DP 的 InferReq 与 KV 引用,
+ # 不能写 shm_infer_released / final token metadata,否则会误把归属 DP
+ # 上尚未结束的请求标记为已完成。所以设置 modify_shm_finish_state 为 False。
+ g_infer_context._filter(
+ finished_request_ids=[req[0] for req in other_dp_reqs],
+ modify_shm_finish_state=False,
+ )
req_ids = [e[0] for e in current_dp_reqs]
@@ -154,8 +161,14 @@ def prefill_normal(
run_reqs_num = len(run_reqs)
with torch.cuda.stream(g_infer_context.get_overlap_stream()):
model_output = self.model.forward(model_input)
+ self._capture_prompt_logprobs_if_needed(model_input, run_reqs, model_output.prompt_logics)
if run_reqs_num > 0:
- _, next_token_ids_cpu, next_token_logprobs_cpu = self._sample_and_scatter_token(
+ (
+ _,
+ next_token_ids_cpu,
+ next_token_logprobs_cpu,
+ next_token_ranks_cpu,
+ ) = self._sample_and_scatter_token(
logits=model_output.logits[:run_reqs_num],
b_req_idx=model_input.b_req_idx[:run_reqs_num],
b_mtp_index=model_input.b_mtp_index[:run_reqs_num],
@@ -183,6 +196,7 @@ def prefill_normal(
run_reqs=run_reqs,
next_token_ids=next_token_ids_cpu,
next_token_logprobs=next_token_logprobs_cpu,
+ next_token_ranks=next_token_ranks_cpu,
run_reqs_update_packs=update_packs,
extra_post_req_handle_func=self.extra_post_req_handle_func,
pd_prefill_chunked_handle_func=self.pd_prefill_chunked_handle_func,
@@ -202,7 +216,12 @@ def decode_normal(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq
with torch.cuda.stream(g_infer_context.get_overlap_stream()):
model_output = self.model.forward(model_input)
if run_reqs_num > 0:
- _, next_token_ids_cpu, next_token_logprobs_cpu = self._sample_and_scatter_token(
+ (
+ _,
+ next_token_ids_cpu,
+ next_token_logprobs_cpu,
+ next_token_ranks_cpu,
+ ) = self._sample_and_scatter_token(
logits=model_output.logits[:run_reqs_num],
b_req_idx=model_input.b_req_idx[:run_reqs_num],
b_mtp_index=model_input.b_mtp_index[:run_reqs_num],
@@ -225,6 +244,7 @@ def decode_normal(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq
run_reqs=run_reqs,
next_token_ids=next_token_ids_cpu,
next_token_logprobs=next_token_logprobs_cpu,
+ next_token_ranks=next_token_ranks_cpu,
run_reqs_update_packs=update_packs,
extra_post_req_handle_func=self.extra_post_req_handle_func,
)
@@ -249,6 +269,8 @@ def prefill_overlap(self, event_pack: OverlapEventPack, prefill_reqs: List[Infer
with torch.cuda.stream(g_infer_context.get_overlap_stream()):
model_output0, model_output1 = self.model.microbatch_overlap_prefill(model_input0, model_input1)
+ self._capture_prompt_logprobs_if_needed(model_input0, run_reqs0, model_output0.prompt_logics)
+ self._capture_prompt_logprobs_if_needed(model_input1, run_reqs1, model_output1.prompt_logics)
logits0 = model_output0.logits
logits1 = model_output1.logits
@@ -266,8 +288,12 @@ def prefill_overlap(self, event_pack: OverlapEventPack, prefill_reqs: List[Infer
b_req_idx = torch.cat((model_input0.b_req_idx[0:req_num0], model_input1.b_req_idx[0:req_num1]), dim=0)
if (req_num0 + req_num1) > 0:
-
- _, next_token_ids_cpu, next_token_logprobs_cpu = self._sample_and_scatter_token(
+ (
+ _,
+ next_token_ids_cpu,
+ next_token_logprobs_cpu,
+ next_token_ranks_cpu,
+ ) = self._sample_and_scatter_token(
logits=logits,
b_req_idx=b_req_idx,
b_mtp_index=b_mtp_index,
@@ -296,6 +322,7 @@ def prefill_overlap(self, event_pack: OverlapEventPack, prefill_reqs: List[Infer
run_reqs=run_reqs,
next_token_ids=next_token_ids_cpu,
next_token_logprobs=next_token_logprobs_cpu,
+ next_token_ranks=next_token_ranks_cpu,
run_reqs_update_packs=update_packs,
extra_post_req_handle_func=self.extra_post_req_handle_func,
pd_prefill_chunked_handle_func=self.pd_prefill_chunked_handle_func,
@@ -335,7 +362,12 @@ def decode_overlap(self, event_pack: OverlapEventPack, decode_reqs: List[InferRe
run_reqs = run_reqs0 + run_reqs1
if (req_num0 + req_num1) > 0:
- _, next_token_ids_cpu, next_token_logprobs_cpu = self._sample_and_scatter_token(
+ (
+ _,
+ next_token_ids_cpu,
+ next_token_logprobs_cpu,
+ next_token_ranks_cpu,
+ ) = self._sample_and_scatter_token(
logits=logits,
b_req_idx=b_req_idx,
b_mtp_index=b_mtp_index,
@@ -358,6 +390,7 @@ def decode_overlap(self, event_pack: OverlapEventPack, decode_reqs: List[InferRe
run_reqs=run_reqs,
next_token_ids=next_token_ids_cpu,
next_token_logprobs=next_token_logprobs_cpu,
+ next_token_ranks=next_token_ranks_cpu,
run_reqs_update_packs=update_packs,
extra_post_req_handle_func=self.extra_post_req_handle_func,
)
@@ -377,13 +410,18 @@ def prefill_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq]
with torch.cuda.stream(g_infer_context.get_overlap_stream()):
model_output: ModelOutput = self.model.forward(model_input)
b_has_out_cpu = model_input.b_prefill_has_output_cpu[0:req_num]
- logits = model_output.logits[0:req_num, :]
+ self._capture_prompt_logprobs_if_needed(model_input, run_reqs, model_output.prompt_logics)
b_req_idx = model_input.b_req_idx[0:req_num]
b_mtp_index = model_input.b_mtp_index[0:req_num]
if req_num > 0:
- next_token_ids, next_token_ids_cpu, next_token_logprobs_cpu = self._sample_and_scatter_token(
- logits=logits,
+ (
+ next_token_ids,
+ next_token_ids_cpu,
+ next_token_logprobs_cpu,
+ next_token_ranks_cpu,
+ ) = self._sample_and_scatter_token(
+ logits=model_output.logits[0:req_num, :],
b_req_idx=b_req_idx,
b_mtp_index=b_mtp_index,
run_reqs=run_reqs,
@@ -421,6 +459,7 @@ def prefill_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq]
run_reqs=run_reqs,
next_token_ids=next_token_ids_cpu,
next_token_logprobs=next_token_logprobs_cpu,
+ next_token_ranks=next_token_ranks_cpu,
run_reqs_update_packs=update_packs,
extra_post_req_handle_func=self.extra_post_req_handle_func,
pd_prefill_chunked_handle_func=self.pd_prefill_chunked_handle_func,
@@ -448,9 +487,12 @@ def decode_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]):
b_req_idx = model_input.b_req_idx[0:req_num]
next_token_ids, next_token_logprobs = sample(logits, run_reqs, self.eos_id)
- next_token_ids_cpu, next_token_logprobs_cpu = self._async_copy_next_token_infos_to_pin_mem(
- next_token_ids, next_token_logprobs
- )
+ next_token_ranks = self._get_next_token_ranks(logits, next_token_ids)
+ (
+ next_token_ids_cpu,
+ next_token_logprobs_cpu,
+ next_token_ranks_cpu,
+ ) = self._async_copy_next_token_infos_to_pin_mem(next_token_ids, next_token_logprobs, next_token_ranks)
# verify the next_token_ids
b_req_mtp_start_loc = [index for index, mtp_index in enumerate(b_mtp_index_cpu) if mtp_index == 0]
@@ -524,6 +566,7 @@ def decode_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]):
run_reqs=verify_ok_reqs,
next_token_ids=next_token_ids_cpu[select_mask],
next_token_logprobs=next_token_logprobs_cpu[select_mask],
+ next_token_ranks=next_token_ranks_cpu[select_mask],
run_reqs_update_packs=update_packs,
extra_post_req_handle_func=self.extra_post_req_handle_func,
)
@@ -652,6 +695,8 @@ def prefill_overlap_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[I
) = padded_overlap_prepare_prefill_inputs(prefill_reqs)
with torch.cuda.stream(g_infer_context.get_overlap_stream()):
model_output0, model_output1 = self.model.microbatch_overlap_prefill(model_input0, model_input1)
+ self._capture_prompt_logprobs_if_needed(model_input0, run_reqs0, model_output0.prompt_logics)
+ self._capture_prompt_logprobs_if_needed(model_input1, run_reqs1, model_output1.prompt_logics)
logits0 = model_output0.logits
logits1 = model_output1.logits
req_num0, req_num1 = len(run_reqs0), len(run_reqs1)
@@ -667,7 +712,12 @@ def prefill_overlap_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[I
b_req_idx = torch.cat((model_input0.b_req_idx[0:req_num0], model_input1.b_req_idx[0:req_num1]), dim=0)
if (req_num0 + req_num1) > 0:
- next_token_ids, next_token_ids_cpu, next_token_logprobs_cpu = self._sample_and_scatter_token(
+ (
+ next_token_ids,
+ next_token_ids_cpu,
+ next_token_logprobs_cpu,
+ next_token_ranks_cpu,
+ ) = self._sample_and_scatter_token(
logits=logits,
run_reqs=run_reqs,
b_req_idx=b_req_idx,
@@ -728,6 +778,7 @@ def prefill_overlap_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[I
run_reqs=run_reqs,
next_token_ids=next_token_ids_cpu,
next_token_logprobs=next_token_logprobs_cpu,
+ next_token_ranks=next_token_ranks_cpu,
run_reqs_update_packs=update_packs,
extra_post_req_handle_func=self.extra_post_req_handle_func,
pd_prefill_chunked_handle_func=self.pd_prefill_chunked_handle_func,
@@ -766,9 +817,12 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf
logits[0:req_num0, :].copy_(logits0[0:req_num0, :], non_blocking=True)
logits[req_num0 : (req_num0 + req_num1), :].copy_(logits1[0:req_num1, :], non_blocking=True)
next_token_ids, next_token_logprobs = sample(logits, run_reqs, self.eos_id)
- next_token_ids_cpu, next_token_logprobs_cpu = self._async_copy_next_token_infos_to_pin_mem(
- next_token_ids, next_token_logprobs
- )
+ next_token_ranks = self._get_next_token_ranks(logits, next_token_ids)
+ (
+ next_token_ids_cpu,
+ next_token_logprobs_cpu,
+ next_token_ranks_cpu,
+ ) = self._async_copy_next_token_infos_to_pin_mem(next_token_ids, next_token_logprobs, next_token_ranks)
b_req_idx = torch.cat((model_input0.b_req_idx[0:req_num0], model_input1.b_req_idx[0:req_num1]), dim=0)
b_mtp_index_cpu = torch.cat((b_mtp_index_cpu0[0:req_num0], b_mtp_index_cpu1[0:req_num1]), dim=0)
@@ -852,6 +906,7 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf
run_reqs=verify_ok_reqs,
next_token_ids=next_token_ids_cpu[select_mask],
next_token_logprobs=next_token_logprobs_cpu[select_mask],
+ next_token_ranks=next_token_ranks_cpu[select_mask],
run_reqs_update_packs=update_packs,
extra_post_req_handle_func=self.extra_post_req_handle_func,
)
diff --git a/lightllm/server/router/model_infer/mode_backend/multi_level_kv_cache.py b/lightllm/server/router/model_infer/mode_backend/multi_level_kv_cache.py
index d0025a03c1..57e4006f4a 100644
--- a/lightllm/server/router/model_infer/mode_backend/multi_level_kv_cache.py
+++ b/lightllm/server/router/model_infer/mode_backend/multi_level_kv_cache.py
@@ -64,6 +64,15 @@ def load_cpu_cache_to_reqs(self, reqs: List[InferReq]):
is_master_in_dp = self.backend.is_master_in_dp
for req in reqs:
page_list = req.shm_req.cpu_cache_match_page_indexes.get_all()
+ # 需要返回 prompt logprobs 的请求不应加载 cpu cache:
+ # 命中后会复用缓存 kv、跳过推理,拿不到对应 logprobs。
+ # match 侧通常已跳过;这里仍要 deref 已 match 的 page,避免引用泄漏。
+ if req.sampling_param.shm_param.prompt_logprobs >= 0:
+ if is_master_in_dp:
+ req.shm_req.cpu_prompt_cache_len = 0
+ all_page_list.extend(page_list)
+ continue
+
page_len_list = req.shm_req.token_hash_page_len_list.get_all()
page_len_start_list = [0] + page_len_list
assert len(page_list) <= len(page_len_list)
diff --git a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py
index 242e1089e7..a8086e63f4 100644
--- a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py
+++ b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py
@@ -66,25 +66,33 @@ def _filter_not_ready_reqs(self, req_ids: List[int]) -> List[InferReq]:
req_obj: InferReq = g_infer_context.requests_mapping[request_id]
if self.is_master_in_dp and req_obj.infer_aborted and req_obj.pd_task_num != 0:
+ # 传输未收尾前每个调度循环都可再投一次,作为幂等补刀
self.info_queue.put(PDAbortReq(request_id=req_obj.req_id, device_id=req_obj.pd_trans_device_id))
if req_obj.pd_task_num != (req_obj.pd_task_failed_num + req_obj.pd_task_success_num):
continue
if req_obj.pd_task_failed_num > 0:
- # 强制停止
+ # KV 传输失败:强制补 finish token 并结束。
+ # abort 优先标 ABORTED;纯传输错误标 ERROR(不再误用 STOP)。
if not req_obj.finish_status.is_finished():
+ finish_status = (
+ FinishStatus.FINISHED_ABORTED if req_obj.infer_aborted else FinishStatus.FINISHED_ERROR
+ )
req_obj.cur_output_len += 1
req_obj.set_next_gen_token_id(next_token_id=0, logprob=0.0, output_len=req_obj.cur_output_len)
- req_obj.finish_status.set_status(FinishStatus.FINISHED_STOP)
+ req_obj.finish_status.set_status(finish_status)
if self.is_master_in_dp:
req_obj.shm_req.shm_cur_output_len = req_obj.cur_output_len
req_obj.shm_req.finish_token_index = req_obj.get_cur_total_len() - 1
- req_obj.shm_req.finish_status.set_status(FinishStatus.FINISHED_STOP)
+ req_obj.shm_req.finish_status.set_status(finish_status)
req_obj.shm_req.candetoken_out_len = req_obj.cur_output_len
- logger.error(f"req_id: {req_obj.req_id} forced to finished, it exits kv transfer error")
+ logger.error(
+ f"req_id: {req_obj.req_id} forced to finished "
+ f"(reason={req_obj.finish_status.get_finish_reason()}), kv transfer error"
+ )
# 提前释放有问题的 mem_index
old_prefix_len = 0 if req_obj.shared_kv_node is None else req_obj.shared_kv_node.node_prefix_total_len
diff --git a/lightllm/server/router/model_infer/mode_backend/rl_backend_ops.py b/lightllm/server/router/model_infer/mode_backend/rl_backend_ops.py
new file mode 100644
index 0000000000..2649c879fc
--- /dev/null
+++ b/lightllm/server/router/model_infer/mode_backend/rl_backend_ops.py
@@ -0,0 +1,343 @@
+import gc
+from typing import List, Optional
+
+import torch
+
+from lightllm.utils.dist_utils import init_custom_process_group
+from lightllm.utils.rl.serialization import LocalSerializedTensor, MultiprocessingSerializer
+from lightllm.utils.rl.tensor_bucket import FlattenedTensorBucket, FlattenedTensorMetadata
+from lightllm.utils.rl.torch_cuda_ipc import cuda_rebuild_device_fallback, monkey_patch_torch_reductions
+from lightllm.utils.torch_memory_saver_utils import MemoryTag
+from lightllm.server.io_struct import (
+ FlushCacheReq,
+ InitWeightsUpdateGroupReq,
+ DestroyWeightsUpdateGroupReq,
+ UpdateWeightsFromDistributedReq,
+ UpdateWeightsFromIPCReq,
+ UpdateWeightsFromTensorReq,
+)
+
+
+class RlBackendOps:
+ MEMORY_TAG_ORDER = (MemoryTag.WEIGHT, MemoryTag.KV_CACHE, MemoryTag.GRAPH)
+
+ SUPPORTED = frozenset(
+ {
+ "flush_cache",
+ "release_memory_occupation",
+ "resume_memory_occupation",
+ "init_weights_update_group",
+ "destroy_weights_update_group",
+ "update_weights_from_distributed",
+ "update_weights_from_tensor",
+ "update_weights_from_ipc",
+ }
+ )
+
+ def __init__(self, backend) -> None:
+ self.backend = backend
+ self._model_update_group = {}
+ self._skip_tensor_updates_reason = None
+ self.logger = backend.logger
+
+ @classmethod
+ def supports(cls, op_name: str) -> bool:
+ return op_name in cls.SUPPORTED
+
+ def dispatch(self, op_name: str, op_args):
+ if not self.supports(op_name):
+ raise ValueError(f"RlBackendOps does not support op {op_name}")
+ return getattr(self, op_name)(op_args)
+
+ def flush_cache(self, request: FlushCacheReq):
+ if self.backend.radix_cache is not None:
+ self.backend.radix_cache.flush_cache()
+ return True, "Succeeded to flush cache."
+
+ def _iter_memory_tags(self, tags: Optional[List[MemoryTag]]):
+ return self.MEMORY_TAG_ORDER if tags is None else tags
+
+ def _clear_cuda_cache(self):
+ torch.cuda.empty_cache()
+ gc.collect()
+
+ def _pause_memory_tags(self, tags: Optional[List[MemoryTag]]):
+ torch.cuda.synchronize()
+ for tag in self._iter_memory_tags(tags):
+ self.backend.model.torch_memory_saver.pause(tag=tag)
+ self._clear_cuda_cache()
+
+ def _resume_memory_tags(self, tags: Optional[List[MemoryTag]]):
+ self._clear_cuda_cache()
+ for tag in self._iter_memory_tags(tags):
+ self.backend.model.torch_memory_saver.resume(tag=tag)
+
+ def release_memory_occupation(self, tags: Optional[List[MemoryTag]]):
+ try:
+ self._pause_memory_tags(tags)
+ self.flush_cache(request=None)
+ return True, "Succeeded to release memory occupation."
+ except Exception as e:
+ self.logger.error(f"release memory occupation failed: {str(e)}")
+ return False, f"release memory occupation failed: {str(e)}"
+
+ def resume_memory_occupation(self, tags: Optional[List[MemoryTag]]):
+ try:
+ self._resume_memory_tags(tags)
+ return True, "Succeeded to resume memory occupation."
+ except Exception as e:
+ self.logger.error(f"resume memory occupation failed: {str(e)}")
+ return False, f"resume memory occupation failed: {str(e)}"
+
+ def init_weights_update_group(self, request: InitWeightsUpdateGroupReq):
+ assert torch.distributed.is_initialized(), "Default torch process group must be initialized"
+
+ assert request.group_name != "", "Group name cannot be empty"
+ rank_offset = request.rank_offset
+ rank = rank_offset + self.backend.rank_in_dp
+ world_size = request.world_size
+ group_name = request.group_name
+ self.logger.info(
+ f"init custom process group: master_address={request.master_address}, master_port={request.master_port}, "
+ f"rank_offset={rank_offset}, rank={rank}, world_size={world_size}, group_name={group_name}, "
+ f" backend={request.backend}"
+ )
+
+ try:
+ if group_name in self._model_update_group:
+ raise ValueError(f"Process group with name {group_name} already exists.")
+
+ self._model_update_group[group_name] = init_custom_process_group(
+ backend=request.backend,
+ init_method=f"tcp://{request.master_address}:{request.master_port}",
+ world_size=world_size,
+ rank=rank,
+ group_name=group_name,
+ )
+ return True, "Succeeded to initialize custom process group."
+
+ except Exception as e:
+ message = f"Failed to initialize custom process group: {e}."
+ self.logger.error(message)
+ return False, message
+
+ def destroy_weights_update_group(self, request: DestroyWeightsUpdateGroupReq):
+ try:
+ if request.group_name in self._model_update_group:
+ pg = self._model_update_group.pop(request.group_name)
+ torch.distributed.destroy_process_group(pg)
+ return True, "Succeeded to destroy custom process group."
+ else:
+ return False, "The group to be destroyed does not exist."
+ except Exception as e:
+ message = f"Failed to destroy custom process group: {e}."
+ self.logger.error(message)
+ return False, message
+
+ def update_weights_from_distributed(self, request: UpdateWeightsFromDistributedReq):
+ """
+ Update model weights online through the custom weight update process group.
+ """
+
+ assert request.group_name in self._model_update_group, (
+ f"Group {request.group_name} not in {list(self._model_update_group.keys())}. "
+ "Please call `init_weights_update_group` first."
+ )
+
+ try:
+ weights = {}
+ handles = []
+ for name, dtype, shape in zip(request.names, request.dtypes, request.shapes):
+ target_dtype = dtype if isinstance(dtype, torch.dtype) else getattr(torch, dtype)
+ weight = torch.empty(shape, dtype=target_dtype, device="cuda")
+ handles.append(
+ torch.distributed.broadcast(
+ weight,
+ src=0,
+ group=self._model_update_group[request.group_name],
+ async_op=True,
+ )
+ )
+ weights[name] = weight
+ for handle in handles:
+ handle.wait()
+
+ self.backend.model.load_weights(weights)
+ return True, "Succeeded to update parameter online from distributed."
+
+ except Exception as e:
+ error_msg = (
+ f"Failed to update parameter online: {e}. "
+ f"The full weights of the ModelRunner are partially updated. "
+ f"Please discard the whole weights."
+ )
+ self.logger.error(error_msg)
+ return False, error_msg
+
+ def _update_weights_from_flattened_bucket(
+ self,
+ flattened_tensor_bucket_dict,
+ ):
+ flattened_tensor = flattened_tensor_bucket_dict["flattened_tensor"]
+ metadata = flattened_tensor_bucket_dict["metadata"]
+
+ converted_metadata = []
+ for meta in metadata:
+ if isinstance(meta, dict):
+ converted_meta = FlattenedTensorMetadata(
+ name=meta["name"],
+ shape=meta["shape"],
+ dtype=meta["dtype"],
+ start_idx=meta["start_idx"],
+ end_idx=meta["end_idx"],
+ numel=meta["numel"],
+ )
+ else:
+ converted_meta = FlattenedTensorMetadata(
+ name=meta.name,
+ shape=meta.shape,
+ dtype=meta.dtype,
+ start_idx=meta.start_idx,
+ end_idx=meta.end_idx,
+ numel=meta.numel,
+ )
+ converted_metadata.append(converted_meta)
+
+ bucket = FlattenedTensorBucket(flattened_tensor=flattened_tensor, metadata=converted_metadata)
+ reconstructed_tensors = bucket.reconstruct_tensors()
+
+ named_tensors = {name: tensor for name, tensor in reconstructed_tensors}
+ loaded, skipped = self._load_compatible_named_tensors(named_tensors)
+
+ return (
+ True,
+ "Succeeded to update parameter online from flattened bucket tensor. "
+ f"loaded={loaded}, skipped={skipped}.",
+ )
+
+ @staticmethod
+ def _iter_named_tensors(named_tensors):
+ if isinstance(named_tensors, dict):
+ return named_tensors.items()
+ return named_tensors
+
+ def _get_tensor_update_skip_reason(self, items):
+ if self._skip_tensor_updates_reason is not None:
+ return self._skip_tensor_updates_reason
+
+ target_config = getattr(self.backend.model, "config", {}) or {}
+ target_is_moe = bool(target_config.get("num_experts") or target_config.get("moe_intermediate_size"))
+ source_has_moe_experts = any(".mlp.experts." in name for name, _ in items)
+ if source_has_moe_experts and not target_is_moe:
+ self._skip_tensor_updates_reason = "received MoE expert weights for a non-MoE backend"
+ self.logger.warning("skip tensor weight updates: %s", self._skip_tensor_updates_reason)
+ return self._skip_tensor_updates_reason
+
+ return None
+
+ def _load_compatible_named_tensors(self, named_tensors):
+ items = list(self._iter_named_tensors(named_tensors))
+ skip_reason = self._get_tensor_update_skip_reason(items)
+ if skip_reason is not None:
+ return 0, len(items)
+
+ def _load_range(weight_items):
+ if not weight_items:
+ return 0, 0
+ weight_dict = dict(weight_items)
+ try:
+ self.backend.model.load_weights(weight_dict)
+ return len(weight_items), 0
+ except Exception as e:
+ if len(weight_items) == 1:
+ name, tensor = weight_items[0]
+ self.logger.warning(
+ "skip incompatible tensor update %s shape=%s dtype=%s: %s",
+ name,
+ tuple(tensor.shape) if hasattr(tensor, "shape") else None,
+ getattr(tensor, "dtype", None),
+ e,
+ )
+ return 0, 1
+
+ split_idx = len(weight_items) // 2
+ left_loaded, left_skipped = _load_range(weight_items[:split_idx])
+ right_loaded, right_skipped = _load_range(weight_items[split_idx:])
+ return left_loaded + right_loaded, left_skipped + right_skipped
+
+ return _load_range(items)
+
+ def update_weights_from_tensor(self, request: UpdateWeightsFromTensorReq):
+ try:
+ monkey_patch_torch_reductions()
+ device_module = torch.get_device_module("cuda")
+ infered_device = device_module.current_device()
+
+ if request.load_format == "flattened_bucket":
+ with cuda_rebuild_device_fallback(infered_device):
+ serialized_named_tensors = MultiprocessingSerializer.deserialize(
+ request.serialized_named_tensors[self.backend.rank_in_dp]
+ )
+ return self._update_weights_from_flattened_bucket(flattened_tensor_bucket_dict=serialized_named_tensors)
+
+ def _unwrap_tensor(tensor, tp_rank, device):
+ if isinstance(tensor, LocalSerializedTensor):
+ tensor = tensor.get(tp_rank)
+ clone = tensor.to(device).clone()
+ del tensor
+ return clone
+
+ with cuda_rebuild_device_fallback(infered_device):
+ named_tensors = MultiprocessingSerializer.deserialize(
+ request.serialized_named_tensors[self.backend.rank_in_dp]
+ )
+ named_tensors = {
+ name: _unwrap_tensor(tensor, tp_rank=self.backend.rank_in_dp, device=infered_device)
+ for name, tensor in self._iter_named_tensors(named_tensors)
+ }
+
+ loaded, skipped = self._load_compatible_named_tensors(named_tensors)
+
+ return True, f"Succeeded to update parameter online from tensor. loaded={loaded}, skipped={skipped}."
+
+ except Exception as e:
+ message = f"Failed to update parameter online from tensor. Reason: {e}."
+ self.logger.error(message)
+
+ return False, message
+
+ def update_weights_from_ipc(self, request: UpdateWeightsFromIPCReq):
+ try:
+ from lightllm.utils.rl.bucketed_weight_transfer import BucketedWeightReceiver, get_zmq_handle
+
+ zmq_handle = request.ipc_handle
+ if isinstance(zmq_handle, dict):
+ zmq_handle = zmq_handle.get(self.backend.rank_in_node, zmq_handle.get(str(self.backend.rank_in_node)))
+ if zmq_handle is None:
+ raise ValueError(f"Missing ipc_handle for rank_in_node={self.backend.rank_in_node}")
+ if zmq_handle in (None, "", "auto"):
+ zmq_handle = get_zmq_handle()
+ use_shm = request.use_shm
+ recv_device = torch.device("cuda", self.backend.current_device_id)
+ self.logger.debug(
+ "[LightLLM] RlBackendOps.update_weights_from_ipc: request.ipc_handle=%r, "
+ "resolved zmq_handle=%r, cuda_device_id=%s",
+ request.ipc_handle,
+ zmq_handle,
+ self.backend.current_device_id,
+ )
+
+ bucketed_weight_receiver = BucketedWeightReceiver(
+ zmq_handle=zmq_handle, device=recv_device, use_shm=use_shm
+ )
+ bucketed_weight_receiver.receive_weights(on_bucket_received=self.backend.model.load_weights)
+ return True, "Succeeded to update parameter online from ipc."
+
+ except Exception as e:
+ import traceback
+
+ traceback.print_exc()
+ message = f"Failed to update parameter online from ipc. Reason: {e}."
+ self.logger.error(message)
+
+ return False, message
diff --git a/lightllm/server/router/model_infer/model_rpc.py b/lightllm/server/router/model_infer/model_rpc.py
index 864a7405b7..5a90a14091 100644
--- a/lightllm/server/router/model_infer/model_rpc.py
+++ b/lightllm/server/router/model_infer/model_rpc.py
@@ -16,7 +16,6 @@
ChunkedPrefillBackend,
FirstTokenConstraintBackend,
OutlinesConstraintBackend,
- ReturnPromptLogProbBackend,
RewardModelBackend,
TokenHealingBackend,
XgrammarBackend,
@@ -28,11 +27,14 @@
PDDPForDecodeNode,
)
from lightllm.server.router.model_infer.mode_backend.redundancy_expert_manager import RedundancyExpertManager
+from lightllm.server.router.model_infer.mode_backend.rl_backend_ops import RlBackendOps
from lightllm.server.core.objs.start_args_type import StartArgs
from lightllm.utils.log_utils import init_logger
from lightllm.utils.graceful_utils import graceful_registry
from lightllm.utils.process_check import start_parent_check_thread
from lightllm.utils.envs_utils import get_unique_server_name
+from lightllm.utils.torch_memory_saver_utils import MemoryTag
+from lightllm.server.io_struct import RlOpReq, RlOpRsp
logger = init_logger(__name__)
@@ -46,6 +48,8 @@ def __init__(self, args, rank: int, rank_in_node: int, node_world_size: int, inf
self.rank = rank
self.rank_in_node = rank_in_node
+ self.backend = None
+ self.rl_backend_ops = None
logger.info(f"Initialized RPC server for rank {self.rank}.")
return
@@ -54,7 +58,6 @@ def exposed_init_model(self, kvargs):
kvargs = obtain(kvargs)
kvargs["rank_id"] = self.rank
self.world_size = kvargs["world_size"]
- return_all_prompt_logprobs = self.args.return_all_prompt_logprobs
use_reward_model = self.args.use_reward_model
diverse_mode = self.args.diverse_mode
is_token_healing = self.args.token_healing_mode
@@ -82,8 +85,6 @@ def exposed_init_model(self, kvargs):
self.backend = DPChunkedPrefillBackend()
elif use_reward_model:
self.backend = RewardModelBackend()
- elif return_all_prompt_logprobs:
- self.backend = ReturnPromptLogProbBackend()
elif diverse_mode:
self.backend = DiversehBackend()
elif is_token_healing:
@@ -99,6 +100,7 @@ def exposed_init_model(self, kvargs):
logger.info(f"use {self.backend.__class__.__name__}")
self.backend.init_model(kvargs)
+ self.rl_backend_ops = RlBackendOps(self.backend) if self.args.enable_rl else None
# only deepseekv3 can support auto_update_redundancy_expert
if self.args.auto_update_redundancy_expert:
@@ -111,6 +113,19 @@ def exposed_init_model(self, kvargs):
def exposed_get_max_total_token_num(self):
return self.backend.get_max_total_token_num()
+ def exposed_rl_op(self, req: RlOpReq) -> RlOpRsp:
+ try:
+ req = obtain(req)
+ if self.rl_backend_ops is None:
+ raise ValueError("RL backend ops is not initialized")
+ if not RlBackendOps.supports(req.op_name):
+ raise ValueError(f"Unsupported RL op {req.op_name}. Supported ops: {sorted(RlBackendOps.SUPPORTED)}")
+ success, ret = self.rl_backend_ops.dispatch(req.op_name, req.op_args)
+ return RlOpRsp(success=success, msg=str(ret), op_name=req.op_name, op_result=ret)
+ except BaseException as e:
+ logger.exception(f"rl op failed: {str(e)}")
+ return RlOpRsp(success=False, msg=f"rl op failed: {str(e)}", op_name=req.op_name)
+
class ModelRpcClient:
def __init__(self, conn):
@@ -134,6 +149,7 @@ async def _func(*args, **kwargs):
self._init_model = async_wrap(self.conn.root.init_model)
self._get_max_total_token_num = async_wrap(self.conn.root.get_max_total_token_num)
+ self._rl_op = async_wrap(self.conn.root.rl_op)
return
async def init_model(self, kvargs):
@@ -145,6 +161,10 @@ async def get_max_total_token_num(self):
ans = self._get_max_total_token_num()
return obtain(await ans)
+ async def rl_op(self, req: RlOpReq) -> RlOpRsp:
+ ans = self._rl_op(req)
+ return obtain(await ans)
+
def _init_env(
args,
@@ -197,7 +217,11 @@ async def start_model_process(
success_event,
),
)
- proc.start()
+ from lightllm.utils.torch_memory_saver_utils import TorchMemorySaverWrapper
+
+ torch_memory_saver = TorchMemorySaverWrapper(args.enable_torch_memory_saver)
+ with torch_memory_saver.configure_subprocess():
+ proc.start()
# Use asyncio.to_thread to make the blocking wait non-blocking
await asyncio.to_thread(success_event.wait, timeout=40)
diff --git a/lightllm/server/router/model_infer/pin_mem_manager.py b/lightllm/server/router/model_infer/pin_mem_manager.py
index 73c7b16b47..d5553e4498 100644
--- a/lightllm/server/router/model_infer/pin_mem_manager.py
+++ b/lightllm/server/router/model_infer/pin_mem_manager.py
@@ -1,7 +1,7 @@
import torch
import threading
import collections
-from typing import List, Dict
+from typing import List, Dict, Union, Sequence
class PinMemTensorManager:
@@ -10,6 +10,9 @@ def __init__(self):
self.key_to_tensor_list: Dict[str, List[torch.Tensor]] = collections.defaultdict(list)
self.key_to_alloc_index: Dict[str, int] = {}
self.buffer_size = 4
+ # 常量 tensor 缓存:逻辑 key -> 已 fill 的 buffer
+ self.key_to_const_cpu_tensor: Dict[str, torch.Tensor] = {}
+ self.key_to_const_gpu_tensor: Dict[str, torch.Tensor] = {}
def alloc_pin_tensor(self, key: str, size: int, dtype: torch.dtype) -> torch.Tensor:
"""
@@ -46,5 +49,57 @@ def async_copy_from_gpu_tensor(self, key: str, gpu_tensor: torch.Tensor) -> torc
pin_mem.copy_(gpu_tensor.view(-1), non_blocking=True)
return pin_mem.view(gpu_tensor.shape)
+ def get_const_cpu_tensor(
+ self,
+ key: str,
+ shape: Sequence[int],
+ fill_value: Union[int, float, bool],
+ dtype: torch.dtype,
+ ) -> torch.Tensor:
+ """返回指定 ``shape`` 的 CPU 常量 tensor 切片(pin_memory,按需扩容)。
+
+ 用途:热路径上需要“占位常量”且不想每 step ``torch.full`` / D2H 时,
+ 例如未开启 ``--enable_rl`` 时 next_token_ranks 固定为 -1。
+ """
+ size = 1
+ for dim in shape:
+ size *= int(dim)
+
+ with self.lock:
+ buf = self.key_to_const_cpu_tensor.get(key)
+ if buf is None or buf.numel() < size:
+ n = max(size, 2048)
+ buf = torch.full((n,), fill_value, dtype=dtype, device="cpu", pin_memory=True)
+ self.key_to_const_cpu_tensor[key] = buf
+ else:
+ assert buf.dtype == dtype, f"const cpu tensor key={key!r} dtype mismatch: {buf.dtype} vs {dtype}"
+ return buf[:size].view(tuple(int(d) for d in shape))
+
+ def get_const_gpu_tensor(
+ self,
+ key: str,
+ shape: Sequence[int],
+ fill_value: Union[int, float, bool],
+ dtype: torch.dtype,
+ ) -> torch.Tensor:
+ """返回指定 ``shape`` 的 GPU 常量 tensor 切片(按需扩容)。
+
+ 与 ``get_const_cpu_tensor`` 对称:热路径上需要 GPU 侧占位常量、又不想每 step
+ ``torch.full`` 时使用。设备取当前 CUDA device。
+ """
+ size = 1
+ for dim in shape:
+ size *= int(dim)
+
+ with self.lock:
+ buf = self.key_to_const_gpu_tensor.get(key)
+ if buf is None or buf.numel() < size:
+ n = max(size, 2048)
+ buf = torch.full((n,), fill_value, dtype=dtype, device="cuda")
+ self.key_to_const_gpu_tensor[key] = buf
+ else:
+ assert buf.dtype == dtype, f"const gpu tensor key={key!r} dtype mismatch: {buf.dtype} vs {dtype}"
+ return buf[:size].view(tuple(int(d) for d in shape))
+
g_pin_mem_manager = PinMemTensorManager()
diff --git a/lightllm/server/router/multinode_tp_helper.py b/lightllm/server/router/multinode_tp_helper.py
new file mode 100644
index 0000000000..e49e1466bc
--- /dev/null
+++ b/lightllm/server/router/multinode_tp_helper.py
@@ -0,0 +1,219 @@
+"""Router 多机纯 TP(nnodes>1 and dp==1)相关逻辑 Mixin。
+
+Public 入口:
+1. ``multinode_tp_generate_new_batch`` — 跨节点调度(内部会调用 abort 阶段1)
+2. ``get_aborted_reqs_from_running_batch_multinode_tp`` — abort 阶段2(running)
+3. ``get_stop_str_matched_reqs_from_running_batch_multinode_tp`` — stop_str 阶段2(running)
+"""
+
+import torch
+import torch.distributed as dist
+
+from typing import List, Optional, Set
+from lightllm.server.router.batch import Batch, Req
+from lightllm.utils.log_utils import init_logger
+
+logger = init_logger(__name__)
+
+
+class RouterMultiNodeTpHelper:
+ """挂到 ``RouterManager``:提供多机 TP 调度与 abort 两阶段接口。"""
+
+ # ==================================================================
+ # Public
+ # ==================================================================
+
+ def multinode_tp_generate_new_batch(self):
+ """跨节点调度入口:barrier → 只调度 → merge → abort 阶段1 → barrier。"""
+ try:
+ dist.barrier(group=self.mulitnode_group)
+ if self.is_multinode_tp_master:
+ new_batch = self._multinode_tp_schedule_as_master()
+ else:
+ new_batch = self._multinode_tp_schedule_as_slave()
+ self.schedule_new_batch = Batch.merge_two_batch(self.schedule_new_batch, new_batch)
+ self._filter_aborted_from_schedule_new_batch()
+ dist.barrier(group=self.mulitnode_group)
+ except Exception as e:
+ logger.exception(str(e))
+ raise e
+ return
+
+ def get_aborted_reqs_from_running_batch_multinode_tp(self) -> List[Req]:
+ """Abort 阶段2:对 running_batch 同步 abort,提取尚未下发 AbortedReqCmd 的请求。"""
+ ans = []
+ running_reqs = [] if self.running_batch is None else self.running_batch.reqs
+ aborted_req_ids = self._broadcast_aborted_req_ids_from_master(running_reqs)
+ if self.is_multinode_tp_slave:
+ for req in running_reqs:
+ if req.request_id in aborted_req_ids:
+ req.is_aborted = True
+
+ for req in running_reqs:
+ if req.is_aborted and req._router_aborted is False:
+ req._router_aborted = True
+ ans.append(req)
+ return ans
+
+ # ==================================================================
+ # Private: 调度
+ # ==================================================================
+
+ def _multinode_tp_schedule_as_master(self) -> Optional[Batch]:
+ """Master:本地 generate_new_batch,广播 req_ids,按就绪标记裁剪;不处理 abort。"""
+ # current_batch: 已占用资源(running + 已调度未推理),供调度估算
+ # new_batch: 本轮从 waiting 新选出的候选
+ current_batch = Batch.merge_two_batch(self.running_batch, self.schedule_new_batch)
+ new_batch = self.req_queue.generate_new_batch(current_batch)
+ req_ids = [req.request_id for req in new_batch.reqs] if new_batch is not None else []
+
+ dist.broadcast_object_list([len(req_ids)], src=0, group=self.mulitnode_group)
+ if len(req_ids) == 0:
+ return None
+
+ dist.broadcast_object_list(req_ids, src=0, group=self.mulitnode_group)
+ # master 本机一定有这些 req;用全 1 参与 MIN all_reduce,确认各 slave waiting 是否已收到
+ select_marks = self._multinode_tp_allreduce_ready_marks([1] * len(req_ids))
+
+ back_req_list = []
+ for req_id, select in zip(req_ids, select_marks):
+ if select == 1:
+ continue
+ # 某节点尚未 ready:打回 waiting,下轮再调度
+ req = new_batch.pop_req(req_id)
+ back_req_list.append(req)
+ self.req_queue.waiting_req_list = back_req_list + self.req_queue.waiting_req_list
+ return None if new_batch.is_clear() else new_batch
+
+ def _multinode_tp_schedule_as_slave(self) -> Optional[Batch]:
+ """Slave:接收 master 的 req_ids,按就绪标记组 batch;不处理 abort。"""
+ req_nums = [None]
+ dist.broadcast_object_list(req_nums, src=0, group=self.mulitnode_group)
+ req_num = req_nums[0]
+ if req_num == 0:
+ return None
+
+ req_ids = [None for _ in range(req_num)]
+ dist.broadcast_object_list(req_ids, src=0, group=self.mulitnode_group)
+
+ id_to_req = {req.request_id: req for req in self.req_queue.waiting_req_list}
+ local_ready = [1 if req_id in id_to_req else 0 for req_id in req_ids]
+ select_marks = self._multinode_tp_allreduce_ready_marks(local_ready)
+
+ select_reqs = []
+ for req_id, select in zip(req_ids, select_marks):
+ if select == 1:
+ select_reqs.append(id_to_req[req_id])
+
+ handled_req_ids = {req.request_id for req in select_reqs}
+ if handled_req_ids:
+ self.req_queue.waiting_req_list = [
+ req for req in self.req_queue.waiting_req_list if req.request_id not in handled_req_ids
+ ]
+
+ if not select_reqs:
+ return None
+ return Batch(-1, reqs=select_reqs, dp_size_in_node=self.dp_size_in_node)
+
+ def _multinode_tp_allreduce_ready_marks(self, local_marks: List[int]) -> List[int]:
+ """对各节点「waiting 是否已有该 req」做 MIN all_reduce:全员 ready 才为 1。"""
+ marks = torch.tensor(local_marks, dtype=torch.int32, device="cpu")
+ dist.all_reduce(marks, op=dist.ReduceOp.MIN, group=self.mulitnode_group)
+ return marks.tolist()
+
+ # ==================================================================
+ # Public: StopStr
+ # ==================================================================
+
+ def get_stop_str_matched_reqs_from_running_batch_multinode_tp(self) -> List[Req]:
+ """同步 stop_str_matched 状态,返回所有节点共同确认的待停止请求。"""
+ running_reqs = [] if self.running_batch is None else list(self.running_batch.reqs)
+ id_to_req = {req.request_id: req for req in running_reqs}
+
+ local_matched_req_ids = [
+ req_id for req_id, req in id_to_req.items() if req.stop_str_matched and not req._router_stop_str_matched
+ ]
+
+ matched_req_ids = self._allgather_stop_str_matched_req_ids(local_matched_req_ids)
+
+ ans = []
+ for req_id in matched_req_ids:
+ req = id_to_req.get(req_id)
+ if req is not None and not req._router_stop_str_matched:
+ req._router_stop_str_matched = True
+ ans.append(req)
+ return ans
+
+ # ==================================================================
+ # Private: Abort
+ # ==================================================================
+
+ def _filter_aborted_from_schedule_new_batch(self):
+ """Abort 阶段1:对 merge 后的 schedule_new_batch 同步 abort,并释放已 abort 请求。"""
+ reqs = [] if self.schedule_new_batch is None else list(self.schedule_new_batch.reqs)
+ # master 用本地 reqs 作为权威源;slave 传空 list,只接收 broadcast 并打标
+ if self.is_multinode_tp_master:
+ aborted_req_ids = self._broadcast_aborted_req_ids_from_master(reqs)
+ else:
+ aborted_req_ids = self._broadcast_aborted_req_ids_from_master([])
+ for req in reqs:
+ if req.request_id in aborted_req_ids:
+ req.is_aborted = True
+
+ if not aborted_req_ids:
+ return
+ if self.schedule_new_batch is None:
+ logger.warning(
+ f"aborted_req_ids non-empty but schedule_new_batch is None, "
+ f"skip release, aborted_ids={sorted(aborted_req_ids)}"
+ )
+ return
+
+ for req_id in aborted_req_ids:
+ if req_id not in self.schedule_new_batch.id_to_reqs:
+ continue
+ req = self.schedule_new_batch.pop_req(req_id)
+ self.req_queue.release_aborted_req(req)
+
+ if self.schedule_new_batch.is_clear():
+ self.schedule_new_batch = None
+ return
+
+ def _broadcast_aborted_req_ids_from_master(self, reqs: List[Req]) -> Set[int]:
+ """以 master 侧 ``reqs`` 中的 ``is_aborted`` 为源,broadcast 到所有节点。"""
+ local_aborted_req_ids = [req.request_id for req in reqs if req.is_aborted]
+ if not self.is_multinode_tp_master:
+ local_aborted_req_ids = []
+
+ aborted_req_num = torch.tensor([len(local_aborted_req_ids)], dtype=torch.int64, device="cpu")
+ dist.broadcast(aborted_req_num, src=0, group=self.mulitnode_group)
+ aborted_req_num = int(aborted_req_num.item())
+ if aborted_req_num == 0:
+ return set()
+
+ if self.is_multinode_tp_master:
+ aborted_req_ids = torch.tensor(local_aborted_req_ids, dtype=torch.int64, device="cpu")
+ else:
+ aborted_req_ids = torch.empty(aborted_req_num, dtype=torch.int64, device="cpu")
+ dist.broadcast(aborted_req_ids, src=0, group=self.mulitnode_group)
+ return {int(req_id) for req_id in aborted_req_ids.tolist()}
+
+ def _allgather_stop_str_matched_req_ids(self, local_matched_req_ids: List[int]) -> List[int]:
+ """各节点独立提供本地 stop_str_matched req_ids,all_gather 后取交集并排序。
+
+ 取交集的原因:
+ - 多机 TP 下同一请求的各 TP shard 可能分布在不同节点;
+ - 只有当一个请求在所有节点的 running_batch 中都 ``stop_str_matched=True`` 时,
+ 才认为该请求真正匹配停止字符串,需要下发 StopStrMatchedReqCmd;
+ - 若某节点未匹配到而其他节点匹配到,说明该请求的 detokenization 状态尚不一致,
+ 应等下一轮调度再确认,避免不一致的下发。
+
+ 排序原因:
+ - 保证各节点返回的 req 顺序一致,避免因集合迭代顺序不同导致下游处理顺序不一致。
+ """
+ gathered = [None] * self.nnodes
+ dist.all_gather_object(gathered, local_matched_req_ids, group=self.mulitnode_group)
+ all_matched = set(gathered[0])
+ for ids in gathered[1:]:
+ all_matched &= set(ids)
+ return sorted(all_matched)
diff --git a/lightllm/server/router/profiler_service.py b/lightllm/server/router/profiler_service.py
index dd27d8d399..c7b8387a6d 100644
--- a/lightllm/server/router/profiler_service.py
+++ b/lightllm/server/router/profiler_service.py
@@ -40,13 +40,15 @@ def start_router_profiler_server(args, profiler_cmd_queue: RouterProfilerCmdQueu
from rpyc.utils.server import ThreadedServer
import lightllm.utils.rpyc_fix_utils as _
+ from lightllm.utils.shm_port_args import get_shm_port_args
+ router_profiler_port = get_shm_port_args().router_profiler_port
server = ThreadedServer(
RouterProfilerService(profiler_cmd_queue),
- port=args.router_profiler_port,
+ port=router_profiler_port,
protocol_config={"allow_pickle": True},
)
thread = threading.Thread(target=server.start, daemon=True)
thread.start()
- logger.info(f"router profiler rpyc server started on port {args.router_profiler_port}")
+ logger.info(f"router profiler rpyc server started on port {router_profiler_port}")
return server, thread
diff --git a/lightllm/server/router/req_queue/base_queue.py b/lightllm/server/router/req_queue/base_queue.py
index 0d1ffe6967..d3b2d7e455 100644
--- a/lightllm/server/router/req_queue/base_queue.py
+++ b/lightllm/server/router/req_queue/base_queue.py
@@ -1,9 +1,13 @@
+import time
from typing import List, Dict
from lightllm.utils.infer_utils import calculate_time
from ..batch import Batch, Req
from lightllm.server.core.objs import FinishStatus
from lightllm.utils.config_utils import get_fixed_kv_len
from lightllm.server.core.objs import StartArgs
+from lightllm.utils.log_utils import init_logger
+
+logger = init_logger(__name__)
class BaseQueue:
@@ -32,6 +36,66 @@ def free_aborted_req_cpu_cache_pages(self, req: Req):
req.cpu_cache_match_page_indexes.clear()
self.router.cpu_cache_client.lock.release()
+ def should_release_aborted_req_in_queue(self, req: Req):
+ # 多节点 TP 的 waiting req abort 状态必须先由 rank 0 broadcast 对齐,
+ # 不能在各节点本地调度队列里提前按各自 shm 状态释放。
+ return req.is_aborted and not self.router.is_multinode_tp
+
+ def mark_aborted_req_finished(self, req: Req):
+ # 未开始推理的请求没有生成 token;这里写入一个 EOS 位置和 aborted 状态,
+ # 让 httpserver recycle loop 能正常结束请求并返回空字符串。
+ input_len = req.input_len
+ req.link_prompt_ids_shm_array()
+ req.link_logprobs_shm_array()
+ req.finish_token_index = input_len
+ req.shm_prompt_ids.arr[input_len] = self.args.eos_id[0]
+ # shm_logprobs 为 structured array: [("logprob", f32), ("rank", i32)]
+ req.shm_logprobs.arr["logprob"][input_len] = 0.0
+ req.shm_logprobs.arr["rank"][input_len] = -1
+ req.finish_status.set_status(FinishStatus.FINISHED_ABORTED)
+
+ # 所有数据准备完后再通知 detokenizer
+ req.candetoken_out_len = 1
+ # 未进 Infer,无 final_token_metadata 可写;置位以免 HTTP 空等
+ req.shm_infer_released = True
+
+ def release_aborted_req(self, req: Req):
+ logger.debug(f"router abort req id {req.request_id} shm_index: {req.index_in_shm_mem}")
+ self.free_aborted_req_cpu_cache_pages(req)
+ self.mark_aborted_req_finished(req)
+ self.router.shm_req_manager.put_back_req_obj(req)
+ return
+
+ def filter_aborted_reqs(self):
+ # 只释放 should_release_aborted_req_in_queue 为真的请求。
+ # 采一波 → sleep 10ms → 再采;前后 request_id 集合完全一致才释放,
+ # 避免同组 abort 标记写全前提前摘掉。多机 TP 下门禁为 False,不进入释放路径。
+ aborted_reqs = [req for req in self.waiting_req_list if self.should_release_aborted_req_in_queue(req)]
+ if not aborted_reqs:
+ return
+
+ prev_ids = {req.request_id for req in aborted_reqs}
+ for _ in range(100):
+ time.sleep(0.01)
+ aborted_reqs = [req for req in self.waiting_req_list if self.should_release_aborted_req_in_queue(req)]
+ cur_ids = {req.request_id for req in aborted_reqs}
+ if prev_ids == cur_ids:
+ break
+ prev_ids = cur_ids
+ else:
+ # 100 次仍未稳定,本轮不释放,下轮调度再试
+ logger.warning(
+ f"aborted reqs not stable after 100 retries, skip release this round, "
+ f"aborted_ids={sorted(prev_ids)}"
+ )
+ return
+
+ aborted_ids = {req.request_id for req in aborted_reqs}
+ self.waiting_req_list = [req for req in self.waiting_req_list if req.request_id not in aborted_ids]
+ for req in aborted_reqs:
+ self.release_aborted_req(req)
+ return
+
def extend(self, req_group: List[Req]):
for req in req_group:
req.sample_params.suggested_dp_index = self.dp_index
diff --git a/lightllm/server/router/req_queue/chunked_prefill/beam_impl.py b/lightllm/server/router/req_queue/chunked_prefill/beam_impl.py
index 63084d9d3b..d5e97317dd 100644
--- a/lightllm/server/router/req_queue/chunked_prefill/beam_impl.py
+++ b/lightllm/server/router/req_queue/chunked_prefill/beam_impl.py
@@ -1,11 +1,7 @@
import uuid
-import time
from typing import List
from ...batch import Batch, Req
from lightllm.server.router.req_queue.base_queue import BaseQueue
-from lightllm.utils.log_utils import init_logger
-
-logger = init_logger(__name__)
class ChunkedBeamContinuesBatchQueue(BaseQueue):
@@ -66,22 +62,6 @@ def _can_add_new_group_reqs(self, cur_handle_group_reqs: List[Req], is_busy, new
else:
return False, new_batch_first_router_need_tokens
- def _filter_aborted_reqs(self):
- # 先移除在等待队列中已经处于aborted状态的请求, 如果发现存在aborted的请求,
- # 则休眠10ms,保证httpserver将所有属于一组的请求都置为aborted请求,再将
- # 请求从队列中移除。
- exist_aborted_req = len([req for req in self.waiting_req_list if req.is_aborted]) > 0
- if exist_aborted_req:
- time.sleep(0.01)
- aborted_reqs = [req for req in self.waiting_req_list if req.is_aborted]
- self.waiting_req_list = [req for req in self.waiting_req_list if not req.is_aborted]
- for req in aborted_reqs:
- req: Req = req
- logger.debug(f"router abort req id {req.request_id} shm_index: {req.index_in_shm_mem}")
- self.free_aborted_req_cpu_cache_pages(req)
- self.router.shm_req_manager.put_back_req_obj(req)
- return
-
# @calculate_time(show=True, min_cost_ms=10)
def generate_new_batch(self, current_batch: Batch):
if len(self.waiting_req_list) == 0:
@@ -93,7 +73,7 @@ def generate_new_batch(self, current_batch: Batch):
if req_is_full:
return None
- self._filter_aborted_reqs()
+ self.filter_aborted_reqs()
if len(self.waiting_req_list) == 0:
return None
diff --git a/lightllm/server/router/req_queue/chunked_prefill/impl.py b/lightllm/server/router/req_queue/chunked_prefill/impl.py
index e82cc7e181..b8bce2a5b5 100644
--- a/lightllm/server/router/req_queue/chunked_prefill/impl.py
+++ b/lightllm/server/router/req_queue/chunked_prefill/impl.py
@@ -2,9 +2,6 @@
import numpy as np
from ...batch import Batch, Req
from lightllm.server.router.req_queue.base_queue import BaseQueue
-from lightllm.utils.log_utils import init_logger
-
-logger = init_logger(__name__)
class ChunkedPrefillQueue(BaseQueue):
@@ -64,6 +61,10 @@ def generate_new_batch(self, current_batch: Batch):
if req_is_full:
return None
+ self.filter_aborted_reqs()
+ if len(self.waiting_req_list) == 0:
+ return None
+
is_busy = self.is_busy()
new_batch_first_router_need_tokens = (
@@ -72,34 +73,23 @@ def generate_new_batch(self, current_batch: Batch):
self._init_cache_list(current_batch, is_busy)
can_run_list = []
- abort_req_list = []
- aborted_count = 0
+ consumed_req_count = 0
waiting_queue = self.waiting_req_list
for req in waiting_queue:
- if req.is_aborted:
- # 由于管理的复杂性,只有没有被调度运行过的请求可以因为abort直接在队列中忽略掉.
- # 暂停的请求需要恢复后,由 router manager 部分来过滤。暂时保持这种处理方法, 否则会导致管理token的泄漏
- aborted_count += 1
- abort_req_list.append(req)
- continue
ok_insert, new_batch_first_router_need_tokens = self._can_add_new_req(
req, is_busy, new_batch_first_router_need_tokens
)
if ok_insert:
+ consumed_req_count += 1
can_run_list.append(req)
else:
break
new_batch = None
if len(can_run_list) != 0:
new_batch = Batch(uuid.uuid4().int, can_run_list, dp_size_in_node=self.dp_size_in_node)
- for req in abort_req_list:
- req: Req = req
- logger.debug(f"router abort req id {req.request_id} shm_index: {req.index_in_shm_mem}")
- self.free_aborted_req_cpu_cache_pages(req)
- self.router.shm_req_manager.put_back_req_obj(req)
- self.waiting_req_list = self.waiting_req_list[len(can_run_list) + aborted_count :]
+ self.waiting_req_list = self.waiting_req_list[consumed_req_count:]
return new_batch
def _calcu_batch_token_load_batch_not_none(self, current_batch: Batch):
diff --git a/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd.py b/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd.py
index 5ec09f5760..f6e33144ef 100644
--- a/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd.py
+++ b/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd.py
@@ -3,9 +3,6 @@
from typing import Tuple
from ...batch import Batch, Req
from lightllm.server.router.req_queue.base_queue import BaseQueue
-from lightllm.utils.log_utils import init_logger
-
-logger = init_logger(__name__)
class PDQueue(BaseQueue):
@@ -64,38 +61,31 @@ def generate_new_batch(self, current_batch: Batch):
if req_is_full:
return None
+ self.filter_aborted_reqs()
+ if len(self.waiting_req_list) == 0:
+ return None
+
estimated_peak_token_num = self._caclu_batch_estimated_peak_token_num(current_batch)
batch_req_num = exist_req_num
can_run_list = []
- abort_req_list = []
- aborted_count = 0
+ consumed_req_count = 0
waiting_queue = self.waiting_req_list
for req in waiting_queue:
- if req.is_aborted:
- # 由于管理的复杂性,只有没有被调度运行过的请求可以因为abort直接在队列中忽略掉.
- # 暂停的请求需要恢复后,由 router manager 部分来过滤。暂时保持这种处理方法, 否则会导致管理token的泄漏
- aborted_count += 1
- abort_req_list.append(req)
- continue
ok_insert, estimated_peak_token_num, batch_req_num = self._can_add_new_req(
req=req, estimated_peak_token_num=estimated_peak_token_num, batch_req_num=batch_req_num
)
if ok_insert:
+ consumed_req_count += 1
can_run_list.append(req)
else:
break
new_batch = None
if len(can_run_list) != 0:
new_batch = Batch(uuid.uuid4().int, can_run_list, dp_size_in_node=self.dp_size_in_node)
- for req in abort_req_list:
- req: Req = req
- logger.debug(f"router abort req id {req.request_id} shm_index: {req.index_in_shm_mem}")
- self.free_aborted_req_cpu_cache_pages(req)
- self.router.shm_req_manager.put_back_req_obj(req)
- self.waiting_req_list = self.waiting_req_list[len(can_run_list) + aborted_count :]
+ self.waiting_req_list = self.waiting_req_list[consumed_req_count:]
return new_batch
def _calcu_batch_token_load_batch_not_none(self, current_batch: Batch):
diff --git a/lightllm/server/router/req_queue/dp_base_queue.py b/lightllm/server/router/req_queue/dp_base_queue.py
index 866e1b9f42..af8f875d4e 100644
--- a/lightllm/server/router/req_queue/dp_base_queue.py
+++ b/lightllm/server/router/req_queue/dp_base_queue.py
@@ -26,6 +26,12 @@ def __init__(self, args, router, base_queue_class, dp_size_in_node) -> None:
self.reqs_waiting_for_dp_index: List[List[Req]] = []
return
+ def release_aborted_req(self, req: Req):
+ dp_index = req.sample_params.suggested_dp_index
+ assert dp_index >= 0 and dp_index < self.dp_size_in_node
+ self.inner_queues[dp_index].release_aborted_req(req)
+ return
+
def get_dp_queue(self, dp_index: int):
assert dp_index < self.dp_size_in_node, "dp index out of range"
return self.inner_queues[dp_index]
diff --git a/lightllm/server/router/rl_rpyc.py b/lightllm/server/router/rl_rpyc.py
new file mode 100644
index 0000000000..e11d03e462
--- /dev/null
+++ b/lightllm/server/router/rl_rpyc.py
@@ -0,0 +1,140 @@
+import asyncio
+import concurrent.futures
+import queue
+import threading
+import rpyc
+import torch
+import torch.distributed as dist
+
+from typing import List, Tuple
+from rpyc.utils.classic import obtain
+from lightllm.server.io_struct import RlOpReq, RlOpRsp
+from lightllm.utils.log_utils import init_logger
+
+logger = init_logger(__name__)
+
+
+class RouterRlOpHelper(rpyc.Service):
+ """Router RL 控制面:rpyc 入口 + 主循环 drain / 多机广播 / 下发 model rpc。
+
+ 挂到 ``RouterManager`` 上;``RouterRlOpQueue`` 经 ``_get_rl_op_queue`` 延迟创建。
+ """
+
+ def exposed_rl_op(self, req: RlOpReq):
+ return self._get_rl_op_queue().submit(obtain(req))
+
+ async def process_rl_ops(self):
+ # 从本地 RL 队列取出本轮待处理的 (req, future)。
+ # 多机 TP 下只有 master 会跑 rpyc service、真正入队;slave 没有提交入口,
+ # 理论上不应依赖本接口拿业务请求(pop 结果恒为空),后续靠 broadcast 对齐 reqs。
+ # 这里仍调用是为了统一 master/slave 代码路径,slave 侧等价于 no-op。
+ pairs = self._get_rl_op_queue().pop_all()
+ reqs: List[RlOpReq] = [req for req, _ in pairs]
+
+ # 多机 TP: master 广播 req;slave 在此处收到同一批 reqs 并跟跑
+ if self.is_multinode_tp:
+ reqs = self._broadcast_rl_ops_to_other_nodes(reqs)
+
+ for i, req in enumerate(reqs):
+ assert isinstance(req, RlOpReq), "rl op request must be RlOpReq"
+ try:
+ ret = await self._rl_op(req)
+ except BaseException as e:
+ logger.exception(f"rl_op failed for {req.op_name}: {e}")
+ ret = RlOpRsp(success=False, msg=f"rl_op error: {e}", op_name=req.op_name)
+ # 多机 TP slave 只跟跑 collective,无权回写 future(仅 master / 单机持有提交方)
+ if self.is_multinode_tp_slave:
+ continue
+ _, fut = pairs[i]
+ if not fut.done():
+ fut.set_result(ret)
+
+ def _get_rl_op_queue(self) -> "RouterRlOpQueue":
+ rl_op_queue = getattr(self, "_rl_op_queue", None)
+ if rl_op_queue is None:
+ self._rl_op_queue = RouterRlOpQueue()
+ return self._rl_op_queue
+
+ def _broadcast_rl_ops_to_other_nodes(self, reqs: List[RlOpReq]):
+ req_num = len(reqs)
+ if self.node_rank == 0:
+ req_nums = [len(reqs)]
+ dist.broadcast_object_list(req_nums, src=0, group=self.mulitnode_group)
+ req_num = req_nums[0]
+ if req_num > 0:
+ dist.broadcast_object_list(reqs, src=0, group=self.mulitnode_group)
+ else:
+ req_nums = [None]
+ dist.broadcast_object_list(req_nums, src=0, group=self.mulitnode_group)
+ req_num = req_nums[0]
+ if req_num > 0:
+ reqs = [None for _ in range(req_num)]
+ dist.broadcast_object_list(reqs, src=0, group=self.mulitnode_group)
+ return reqs
+
+ async def _rl_op(self, req: RlOpReq) -> RlOpRsp:
+ rl_op_tasks = []
+ for model_rpc_client in self.model_rpc_clients:
+ rl_op_tasks.append(model_rpc_client.rl_op(req))
+ all_ret = await asyncio.gather(*rl_op_tasks)
+ # 优先返回第一个失败结果(带具体错误信息);全部成功则用第一个
+ ret: RlOpRsp = all_ret[0]
+ for res in all_ret:
+ if not res.success:
+ ret = res
+ break
+ ret.success = all(res.success for res in all_ret)
+
+ if self.is_multinode_tp:
+ # True/False -> 1/0;MIN all_reduce:任一节点失败则全体 success=False
+ success_flag = torch.tensor([1 if ret.success else 0], dtype=torch.int32, device="cpu")
+ dist.all_reduce(success_flag, op=dist.ReduceOp.MIN, group=self.mulitnode_group)
+ ret.success = success_flag.item() == 1
+ return ret
+
+
+class RouterRlOpQueue:
+ def __init__(self):
+ self._queue: "queue.Queue[Tuple[RlOpReq, concurrent.futures.Future]]" = queue.Queue()
+
+ def submit(self, req: RlOpReq, timeout: float = 300.0) -> RlOpRsp:
+ fut: concurrent.futures.Future = concurrent.futures.Future()
+ self._queue.put((req, fut))
+ try:
+ return fut.result(timeout=timeout)
+ except concurrent.futures.TimeoutError:
+ return RlOpRsp(
+ success=False,
+ msg=f"rl op {req.op_name} timeout after {timeout}s",
+ op_name=req.op_name,
+ )
+
+ def pop_all(self) -> List[Tuple[RlOpReq, concurrent.futures.Future]]:
+ pairs = []
+ while True:
+ try:
+ pairs.append(self._queue.get_nowait())
+ except queue.Empty:
+ break
+ return pairs
+
+
+def start_router_rl_rpyc_server(args, router: RouterRlOpHelper):
+ if args.node_rank != 0:
+ return None, None
+
+ from rpyc.utils.server import ThreadedServer
+ import lightllm.utils.rpyc_fix_utils as _
+ from lightllm.utils.shm_port_args import get_shm_port_args
+
+ rl_rpyc_port = get_shm_port_args().rl_rpyc_port
+ server = ThreadedServer(
+ router,
+ hostname="127.0.0.1",
+ port=rl_rpyc_port,
+ protocol_config={"allow_pickle": True, "sync_request_timeout": 600},
+ )
+ thread = threading.Thread(target=server.start, name="rl_rpyc_server", daemon=True)
+ thread.start()
+ logger.info(f"router rl rpyc server started on port {rl_rpyc_port}")
+ return server, thread
diff --git a/lightllm/server/visualserver/manager.py b/lightllm/server/visualserver/manager.py
index 1dffdaf681..873eff8456 100644
--- a/lightllm/server/visualserver/manager.py
+++ b/lightllm/server/visualserver/manager.py
@@ -21,6 +21,7 @@
from lightllm.utils.graceful_utils import graceful_registry
from lightllm.utils.process_check import start_parent_check_thread
from lightllm.utils.envs_utils import get_unique_server_name
+from lightllm.utils.shm_port_args import get_shm_port_args
from rpyc.utils.classic import obtain
@@ -33,22 +34,23 @@ def __init__(
args: StartArgs,
):
self.args = args
+ ports = get_shm_port_args()
context = zmq.Context(2)
enable_audio = not args.disable_audio
if enable_audio:
self.send_to_next_module = context.socket(zmq.PUSH)
- self.send_to_next_module.connect(f"{args.zmq_mode}127.0.0.1:{args.audio_port}")
+ self.send_to_next_module.connect(f"{args.zmq_mode}127.0.0.1:{ports.audio_port}")
else:
if args.enable_cpu_cache:
self.send_to_next_module = context.socket(zmq.PUSH)
- self.send_to_next_module.connect(f"{args.zmq_mode}127.0.0.1:{args.multi_level_kv_cache_port}")
+ self.send_to_next_module.connect(f"{args.zmq_mode}127.0.0.1:{ports.multi_level_kv_cache_port}")
else:
self.send_to_next_module = context.socket(zmq.PUSH)
- self.send_to_next_module.connect(f"{args.zmq_mode}127.0.0.1:{args.router_port}")
+ self.send_to_next_module.connect(f"{args.zmq_mode}127.0.0.1:{ports.router_port}")
self.zmq_recv_socket = context.socket(zmq.PULL)
- self.zmq_recv_socket.bind(f"{args.zmq_mode}127.0.0.1:{args.visual_port}")
- self.cache_client = rpyc.connect("localhost", args.cache_port, config={"allow_pickle": True})
+ self.zmq_recv_socket.bind(f"{args.zmq_mode}127.0.0.1:{ports.visual_port}")
+ self.cache_client = rpyc.connect("localhost", ports.cache_port, config={"allow_pickle": True})
self.cache_client._channel.stream.sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1)
self.model_weightdir = args.model_dir
self.vit_dp = args.visual_dp
@@ -70,6 +72,8 @@ async def wait_to_model_ready(self):
self.model_rpcs[dp_rank_id].append(rpc_model)
init_model_ret = []
+ ports = get_shm_port_args()
+ visual_nccl_ports = ports.visual_nccl_ports
for dp_rank_id in range(self.vit_dp): # async init model process
for tp_rank_id in range(self.vit_tp):
device_id = self.args.visual_gpu_ids[dp_rank_id * self.vit_tp + tp_rank_id]
@@ -77,11 +81,11 @@ async def wait_to_model_ready(self):
"weight_dir": self.model_weightdir,
"device_id": device_id,
"vit_tp": self.vit_tp,
- "cache_port": self.args.cache_port,
+ "cache_port": ports.cache_port,
"tp_rank_id": tp_rank_id,
"dp_rank_id": dp_rank_id,
"data_type": self.args.data_type,
- "visual_nccl_port": self.args.visual_nccl_ports[dp_rank_id],
+ "visual_nccl_port": visual_nccl_ports[dp_rank_id],
"quant_type": self.args.vit_quant_type,
"quant_cfg": self.args.vit_quant_cfg,
"max_batch_size": max(self.infer_batch_size // self.vit_dp, 1),
diff --git a/lightllm/server/visualserver/model_infer/__init__.py b/lightllm/server/visualserver/model_infer/__init__.py
index 3e74793634..b08e2d13e2 100644
--- a/lightllm/server/visualserver/model_infer/__init__.py
+++ b/lightllm/server/visualserver/model_infer/__init__.py
@@ -11,6 +11,7 @@
from rpyc.utils.server import ThreadedServer
from lightllm.utils.graceful_utils import graceful_registry
from lightllm.utils.envs_utils import get_env_start_args, get_unique_server_name
+from lightllm.utils.process_check import start_parent_check_thread
from .model_rpc_client import VisualModelRpcClient
from .model_rpc import VisualModelRpcServer
from ..objs import rpyc_config
@@ -20,6 +21,7 @@ def _init_env(socket_path: str, success_event):
# 注册graceful 退出的处理
graceful_registry(inspect.currentframe().f_code.co_name)
setproctitle.setproctitle(f"lightllm::{get_unique_server_name()}::visual_model_infer")
+ start_parent_check_thread()
import lightllm.utils.rpyc_fix_utils as _
diff --git a/lightllm/server/visualserver/model_infer/model_rpc.py b/lightllm/server/visualserver/model_infer/model_rpc.py
index 6e0f842e63..68e0a97ca1 100644
--- a/lightllm/server/visualserver/model_infer/model_rpc.py
+++ b/lightllm/server/visualserver/model_infer/model_rpc.py
@@ -24,6 +24,7 @@
from lightllm.utils.infer_utils import set_random_seed
from lightllm.utils.dist_utils import init_vision_distributed_env
from lightllm.utils.envs_utils import get_env_start_args
+from lightllm.utils.shm_port_args import get_shm_port_args
from lightllm.server.embed_cache.embed_cache_client import CpuEmbedCacheClient
from lightllm.server.visualserver import set_vit_att_backend
from lightllm.server.embed_cache.afs_utils import SepEmbedHandler
@@ -128,7 +129,7 @@ def exposed_init_model(self, kvargs):
self.afs_handler = SepEmbedHandler(
afs_embed_dir=self.args.afs_image_embed_dir,
redis_host=self.args.config_server_host,
- redis_port=self.args.config_server_visual_redis_port,
+ redis_port=get_shm_port_args().config_server_visual_redis_port,
capacity=self.args.afs_embed_capacity,
)
diff --git a/lightllm/server/visualserver/proxy_manager.py b/lightllm/server/visualserver/proxy_manager.py
index 0c977b2aa9..f5b7d58acf 100644
--- a/lightllm/server/visualserver/proxy_manager.py
+++ b/lightllm/server/visualserver/proxy_manager.py
@@ -23,6 +23,7 @@
from lightllm.utils.graceful_utils import graceful_registry
from lightllm.utils.process_check import start_parent_check_thread
from lightllm.utils.envs_utils import get_unique_server_name
+from lightllm.utils.shm_port_args import get_shm_port_args
from rpyc.utils.classic import obtain
from lightllm.server.embed_cache.utils import read_shm, get_shm_name_data
from .manager import VisualManager
@@ -43,10 +44,11 @@ def __init__(
self.cpu_embed_cache_client = CpuEmbedCacheClient(create_meta_data=False, init_shm_data=False, pin_shm=False)
+ ports = get_shm_port_args()
self.afs_handler = SepEmbedHandler(
afs_embed_dir=self.args.afs_image_embed_dir,
redis_host=self.args.config_server_host,
- redis_port=self.args.config_server_visual_redis_port,
+ redis_port=ports.config_server_visual_redis_port,
capacity=self.args.afs_embed_capacity,
)
@@ -140,7 +142,8 @@ async def loop_to_connect_remote_visual_server(self):
counter = 0
error_counter = 0
while True:
- uri = f"http://{self.args.config_server_host}:{self.args.config_server_port}/registered_visual_objects"
+ config_server_port = get_shm_port_args().config_server_port
+ uri = f"http://{self.args.config_server_host}:{config_server_port}/registered_visual_objects"
try:
async with httpx.AsyncClient(timeout=10.0) as client:
response = await client.get(uri)
diff --git a/lightllm/server/visualserver/visual_only_manager.py b/lightllm/server/visualserver/visual_only_manager.py
index b06713d87c..835a25fe3e 100644
--- a/lightllm/server/visualserver/visual_only_manager.py
+++ b/lightllm/server/visualserver/visual_only_manager.py
@@ -25,6 +25,7 @@
from lightllm.utils.graceful_utils import graceful_registry
from lightllm.utils.process_check import start_parent_check_thread
from lightllm.utils.envs_utils import get_unique_server_name
+from lightllm.utils.shm_port_args import get_shm_port_args
from rpyc.utils.classic import obtain
from lightllm.server.embed_cache.utils import create_shm, get_shm_name_data, free_shm
from .manager import VisualManager
@@ -68,15 +69,16 @@ async def register_to_config_server_loop(self, args: StartArgs):
else:
host_ip = args.host
+ ports = get_shm_port_args()
while True:
try:
- uri = f"ws://{args.config_server_host}:{args.config_server_port}/visual_register"
+ uri = f"ws://{args.config_server_host}:{ports.config_server_port}/visual_register"
async with websockets.connect(uri, max_queue=(2048 * 1024, 2048 * 1023)) as websocket:
sock = websocket.transport.get_extra_info("socket")
sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1)
- vit_obj = VIT_Obj(node_id=args.visual_node_id, host_ip=host_ip, port=args.visual_rpyc_port)
+ vit_obj = VIT_Obj(node_id=args.visual_node_id, host_ip=host_ip, port=ports.visual_rpyc_port)
await websocket.send(pickle.dumps(vit_obj))
logger.info(f"Sent registration vit_obj: {vit_obj}")
@@ -102,6 +104,7 @@ async def wait_to_model_ready(self):
self.model_rpcs[dp_rank_id].append(rpc_model)
init_model_ret = []
+ visual_nccl_ports = get_shm_port_args().visual_nccl_ports
for dp_rank_id in range(self.vit_dp): # async init model process
for tp_rank_id in range(self.vit_tp):
device_id = self.args.visual_gpu_ids[dp_rank_id * self.vit_tp + tp_rank_id]
@@ -113,7 +116,7 @@ async def wait_to_model_ready(self):
"tp_rank_id": tp_rank_id,
"dp_rank_id": dp_rank_id,
"data_type": self.args.data_type,
- "visual_nccl_port": self.args.visual_nccl_ports[dp_rank_id],
+ "visual_nccl_port": visual_nccl_ports[dp_rank_id],
"quant_type": self.args.vit_quant_type,
"quant_cfg": self.args.vit_quant_cfg,
"max_batch_size": max(self.infer_batch_size // self.vit_dp, 1),
@@ -197,7 +200,7 @@ def handle_exception(loop, context):
from .objs import rpyc_config
- t = rpyc.ThreadedServer(visualserver, port=args.visual_rpyc_port, protocol_config=rpyc_config)
+ t = rpyc.ThreadedServer(visualserver, port=get_shm_port_args().visual_rpyc_port, protocol_config=rpyc_config)
except Exception as e:
logger.exception(str(e))
visualserver.clean_up()
diff --git a/lightllm/utils/dist_utils.py b/lightllm/utils/dist_utils.py
index 5b9705ed0e..c3f4684c94 100644
--- a/lightllm/utils/dist_utils.py
+++ b/lightllm/utils/dist_utils.py
@@ -80,12 +80,15 @@ def init_vision_distributed_env(kvargs):
device_id = kvargs["device_id"]
set_current_device_id(device_id)
torch.cuda.set_device(device_id)
+ # 不要在init_process_group时,显示的传入device_id
+ # 这会触发torch的device-bound split优化,会默认后面想加入新进程组的rank
+ # 都已经存在于默认组,这样RL更新weight的init_group时,外部想加入的组,在执行
+ # 通信原语时例如all_reduce,会永远等不到LightLLM默认组里的回复,从而导致错误结果。
dist.init_process_group(
"nccl",
init_method=f'tcp://127.0.0.1:{kvargs["visual_nccl_port"]}',
rank=kvargs["tp_rank_id"],
world_size=tp_world_size,
- device_id=torch.device(f"cuda:{device_id}"),
)
# warmup nccl communicator
_a = torch.zeros([1]).to(f"cuda:{device_id}")
@@ -150,7 +153,6 @@ def init_distributed_env(kvargs):
init_method=f'tcp://{kvargs["nccl_host"]}:{kvargs["nccl_port"]}',
rank=kvargs["rank_id"],
world_size=kvargs["world_size"],
- device_id=torch.device(f"cuda:{device_id}"),
)
# warmup nccl communicator
_a = torch.zeros([1]).to(f"cuda:{device_id}")
@@ -290,6 +292,7 @@ def create_dp_special_inter_group(backend):
def _init_nccl_env():
from lightllm.utils.envs_utils import get_env_start_args
+ from lightllm.utils.shm_port_args import get_shm_port_args
args = get_env_start_args()
@@ -298,8 +301,9 @@ def _init_nccl_env():
os.environ["TORCHELASTIC_USE_AGENT_STORE"] = "True"
rank_id = get_global_rank()
world_size = get_global_world_size()
- ip_port = f"{args.config_server_host}:{args.config_server_port}"
- params = f"tcp_store_port={args.nccl_port}&&rank_id={rank_id}&&world_size={world_size}"
+ ports = get_shm_port_args()
+ ip_port = f"{args.config_server_host}:{ports.config_server_port}"
+ params = f"tcp_store_port={ports.nccl_port}&&rank_id={rank_id}&&world_size={world_size}"
if rank_id == 0:
# 当使用外部config server 启动的tcpStore来初始化nccl时,需要保证配置了config_server_host.
@@ -316,3 +320,71 @@ def _init_nccl_env():
assert response.status_code == 200, f"Failed to init config server nccl tcp store: {response.status_code}"
return
+
+
+# copy from https://github.com/sgl-project/sglang/blob/main/python/sglang/srt/utils/common.py#L1675
+def init_custom_process_group(
+ backend=None,
+ init_method=None,
+ timeout=None,
+ world_size=-1,
+ rank=-1,
+ store=None,
+ group_name=None,
+ pg_options=None,
+ device_id=None,
+):
+ from torch.distributed.distributed_c10d import (
+ Backend,
+ PrefixStore,
+ _new_process_group_helper,
+ _world,
+ default_pg_timeout,
+ rendezvous,
+ )
+
+ assert (store is None) or (init_method is None), "Cannot specify both init_method and store."
+
+ if store is not None:
+ assert world_size > 0, "world_size must be positive if using store"
+ assert rank >= 0, "rank must be non-negative if using store"
+ elif init_method is None:
+ init_method = "env://"
+
+ if backend:
+ backend = Backend(backend)
+ else:
+ backend = Backend("undefined")
+
+ if timeout is None:
+ timeout = default_pg_timeout
+
+ # backward compatible API
+ if store is None:
+ rendezvous_iterator = rendezvous(init_method, rank, world_size, timeout=timeout)
+ store, rank, world_size = next(rendezvous_iterator)
+ store.set_timeout(timeout)
+
+ # Use a PrefixStore to avoid accidental overrides of keys used by
+ # different systems (e.g. RPC) in case the store is multi-tenant.
+ store = PrefixStore(group_name, store)
+
+ # NOTE: The pg_options parameter was renamed into backend_options in PyTorch 2.6.0
+ # https://github.com/pytorch/pytorch/commit/a0c7029a75628cd5fa8df83c0de0ea98ee7fd844
+ # We need to determine the appropriate parameter name based on PyTorch version
+ pg_options_param_name = "backend_options" if str(torch.__version__) >= "2.6" else "pg_options"
+ pg, _ = _new_process_group_helper(
+ world_size,
+ rank,
+ [],
+ backend,
+ store,
+ group_name=group_name,
+ **{pg_options_param_name: pg_options},
+ timeout=timeout,
+ device_id=device_id,
+ )
+
+ _world.pg_group_ranks[pg] = {i: i for i in range(world_size)}
+
+ return pg
diff --git a/lightllm/utils/multinode_utils.py b/lightllm/utils/multinode_utils.py
index ffe3c1208c..bf2e0aba8b 100644
--- a/lightllm/utils/multinode_utils.py
+++ b/lightllm/utils/multinode_utils.py
@@ -1,6 +1,7 @@
import zmq
import socket
from lightllm.utils.log_utils import init_logger
+from lightllm.utils.shm_port_args import get_shm_port_args
logger = init_logger(__name__)
@@ -11,14 +12,15 @@ def send_and_receive_node_ip(args):
# 一些通信组件转发请求信息给从节点。
is_multinode_tp = args.dp == 1 and args.nnodes > 1
if is_multinode_tp:
+ base_port = get_shm_port_args().multinode_httpmanager_port
if args.node_rank == 0:
args.child_ips = None
args.child_ips = []
for i in range(1, args.nnodes):
context = zmq.Context(2)
comm_socket = context.socket(zmq.PULL)
- comm_socket.bind(f"tcp://*:{args.multinode_httpmanager_port + i + 100}")
- logger.info(f"binding port {args.multinode_httpmanager_port + i + 100}")
+ comm_socket.bind(f"tcp://*:{base_port + i + 100}")
+ logger.info(f"binding port {base_port + i + 100}")
args.child_ips.append(comm_socket.recv_pyobj())
comm_socket.close()
logger.info(f"Received child IPs: {args.child_ips}")
@@ -26,7 +28,7 @@ def send_and_receive_node_ip(args):
local_ip = socket.gethostbyname(socket.gethostname())
context = zmq.Context(2)
comm_socket = context.socket(zmq.PUSH)
- comm_socket.connect(f"tcp://{args.nccl_host}:{args.multinode_httpmanager_port + args.node_rank + 100}")
- logger.info(f"connecting to {args.nccl_host}:{args.multinode_httpmanager_port + args.node_rank + 100}")
+ comm_socket.connect(f"tcp://{args.nccl_host}:{base_port + args.node_rank + 100}")
+ logger.info(f"connecting to {args.nccl_host}:{base_port + args.node_rank + 100}")
comm_socket.send_pyobj(local_ip)
comm_socket.close()
diff --git a/lightllm/utils/net_utils.py b/lightllm/utils/net_utils.py
index b87096d945..16ae239f62 100644
--- a/lightllm/utils/net_utils.py
+++ b/lightllm/utils/net_utils.py
@@ -1,48 +1,11 @@
import socket
import subprocess
import ipaddress
-import random
from lightllm.utils.log_utils import init_logger
logger = init_logger(__name__)
-def alloc_can_use_network_port(num=3, used_ports=None, from_port_num=10000):
- port_list = []
- for port in range(from_port_num, 65536):
- with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
- result = s.connect_ex(("localhost", port))
- if result != 0 and port not in used_ports:
- port_list.append(port)
- if len(port_list) > num * 30:
- break
-
- if len(port_list) < num:
- return None
-
- random.shuffle(port_list)
- return port_list[0:num]
-
-
-def alloc_can_use_port(min_port, max_port):
- port_list = []
- for port in range(min_port, max_port):
- with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
- result = s.connect_ex(("localhost", port))
- if result != 0:
- port_list.append(port)
- return port_list
-
-
-def find_available_port(start_port, end_port):
- for port in range(start_port, end_port + 1):
- with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as sock:
- result = sock.connect_ex(("localhost", port))
- if result != 0:
- return port
- return None
-
-
def get_hostname_ip():
try:
result = subprocess.run(["hostname", "-i"], capture_output=True, text=True, check=True)
@@ -63,22 +26,15 @@ def is_valid_ipv6_address(address: str) -> bool:
return False
-class PortLocker:
- def __init__(self, ports):
- self.ports = ports
- self.sockets = [socket.socket(socket.AF_INET, socket.SOCK_STREAM) for _ in range(len(self.ports))]
- for _socket in self.sockets:
- _socket.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
+def validate_ports(ports: list):
+ """校验端口列表:列表内不重复,且当前均可 bind。"""
+ if len(ports) != len(set(ports)):
+ raise RuntimeError(f"conflicting ports in list: {ports}")
- def lock_port(self):
- for _socket, _port in zip(self.sockets, self.ports):
+ for port in ports:
+ with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as sock:
try:
- _socket.bind(("", _port))
- _socket.listen(1)
- except Exception as e:
- logger.error(f"port {_port} has been used")
- raise e
-
- def release_port(self):
- for _socket in self.sockets:
- _socket.close()
+ sock.bind(("", port))
+ except OSError as e:
+ raise RuntimeError(f"port {port} has been used") from e
+ logger.info(f"validated ports: {ports}")
diff --git a/lightllm/utils/redis_utils.py b/lightllm/utils/redis_utils.py
index 30b4ae6450..cfff577f30 100644
--- a/lightllm/utils/redis_utils.py
+++ b/lightllm/utils/redis_utils.py
@@ -1,5 +1,6 @@
import subprocess
from lightllm.utils.log_utils import init_logger
+from lightllm.utils.shm_port_args import get_shm_port_args
logger = init_logger(__name__)
@@ -8,7 +9,7 @@ def start_redis_service(args):
"""launch redis service"""
config_server_host = args.config_server_host
- redis_port = args.config_server_visual_redis_port
+ redis_port = get_shm_port_args().config_server_visual_redis_port
try:
subprocess.run(
["redis-cli", "-h", config_server_host, "-p", str(redis_port), "FLUSHALL", "ASYNC"], check=False, timeout=2
diff --git a/lightllm/utils/rl/__init__.py b/lightllm/utils/rl/__init__.py
new file mode 100644
index 0000000000..12d5ab89e5
--- /dev/null
+++ b/lightllm/utils/rl/__init__.py
@@ -0,0 +1,6 @@
+"""
+Utilities used by RL weight-update and colocated rollout integration paths.
+
+This package groups the CUDA IPC serializer, tensor bucketing helpers, and
+bucketed weight-transfer protocol used by the RL endpoints.
+"""
diff --git a/lightllm/utils/rl/bucketed_weight_transfer.py b/lightllm/utils/rl/bucketed_weight_transfer.py
new file mode 100644
index 0000000000..4849497f1a
--- /dev/null
+++ b/lightllm/utils/rl/bucketed_weight_transfer.py
@@ -0,0 +1,302 @@
+# Copyright 2025 Bytedance Ltd. and/or its affiliates
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+"""
+Bucketed RL weight transfer via ZMQ plus CUDA IPC or shared memory fallback.
+
+This module builds on torch_cuda_ipc.py for CUDA IPC device handling. It owns
+the higher-level transfer protocol: the sender publishes one reusable
+communication buffer, and the receiver rebuilds that buffer on its target
+device before applying each metadata-described bucket to model weights.
+
+Copied from:
+https://github.com/verl-project/verl/blob/main/verl/workers/rollout/vllm_rollout/bucketed_weight_transfer.py
+"""
+
+import gc
+from multiprocessing import shared_memory
+from typing import TypedDict
+
+import torch
+import zmq
+from torch.multiprocessing.reductions import reduce_tensor
+from lightllm.utils.rl.torch_cuda_ipc import (
+ cuda_device_to_uuid,
+ get_current_device_id,
+ get_current_device_name,
+ rebuild_cuda_ipc_tensor,
+)
+
+
+def get_zmq_handle() -> str:
+ return f"ipc:///tmp/rl-colocate-zmq-{cuda_device_to_uuid(get_current_device_id())}.sock"
+
+
+class TensorMetadata(TypedDict):
+ name: str
+ shape: torch.Size
+ dtype: torch.dtype
+ offset: int
+
+
+def create_shared_memory(size: int, name: str):
+ """Create shared memory for weight transfer. If already exists, attach to it."""
+ try:
+ shm = shared_memory.SharedMemory(name=name, create=True, size=size)
+ except FileExistsError:
+ shm = shared_memory.SharedMemory(name=name)
+ assert shm.size >= size, f"Stale shm segment '{name}': expected {size} bytes, got {shm.size}"
+ return shm
+
+
+def rebuild_shared_memory(name: str, size: int, dtype=torch.uint8):
+ """Rebuild tensor from shared memory."""
+ shm = shared_memory.SharedMemory(name=name)
+ tensor = torch.frombuffer(shm.buf[:size], dtype=dtype)
+
+ return tensor, shm
+
+
+class BucketedWeightSender:
+ """
+ Send model weights via bucketed IPC transfer over ZMQ.
+
+ Packs weight tensors into a fixed-size communication buffer and sends them
+ in buckets to the receiver. Supports CUDA IPC and shared memory fallback.
+
+ Args:
+ zmq_handle: ZMQ IPC socket path (e.g., "ipc:///tmp/rl-colocate-zmq-.sock")
+ bucket_size_mb: Communication buffer size in MB
+ use_shm: Use shared memory instead of CUDA IPC (for NPU compatibility)
+ """
+
+ def __init__(
+ self,
+ zmq_handle: str,
+ bucket_size_mb: int = 512,
+ use_shm: bool = False,
+ ):
+ self.zmq_handle = zmq_handle
+ self.bucket_size_mb = bucket_size_mb
+ self.bucket_size = int(bucket_size_mb) << 20
+ self.use_shm = use_shm
+
+ self.zmq_context = zmq.Context.instance()
+ self.socket = None
+ self.buffer = None
+ self.shm = None
+
+ async def async_send_weights(self, weights):
+ """
+ Send weights to the receiver. Accepts a sync generator or async iterator.
+
+ Args:
+ weights: Generator or async iterator yielding (name, tensor) pairs
+ """
+ from verl.workers.rollout.utils import ensure_async_iterator
+
+ try:
+ self._init_socket()
+ self._init_buffer()
+
+ # send bucket weights
+ offset = 0
+ bucket_meta: dict[str, TensorMetadata] = {}
+ # dtype = PrecisionType.to_dtype(self.config.dtype)
+ async for name, weight in ensure_async_iterator(weights):
+ # model parameters are in fp32 full precision
+ # (vermouth1992) we should not force cast weight here because some parameters
+ # (such as moe gate) have to keep fp32 precision. If a weight is bf16 in the rollout side,
+ # the rollout should automatically cast on demand. However, this would incur a higher weight
+ # transfer volume.
+ # weight = weight.to(dtype, non_blocking=True)
+
+ # fill the tensor bucket
+ if offset + weight.nbytes > self.bucket_size:
+ torch.cuda.synchronize()
+ self.socket.send_pyobj({"bucket_meta": bucket_meta, "is_last": False})
+ self.socket.recv()
+ bucket_meta = {}
+ offset = 0
+
+ # TODO: slice embedding layer weight into chunks
+ assert offset + weight.nbytes <= self.bucket_size, (
+ f"Weight {name}({weight.shape}, {weight.dtype}) is too large to fit in the bucket."
+ f"Please increase rollout.update_weights_bucket_megabytes({self.bucket_size_mb} MB)."
+ )
+ bucket_meta[name] = {
+ "name": name,
+ "shape": weight.shape,
+ "dtype": weight.dtype,
+ "offset": offset,
+ }
+ self.buffer[offset : offset + weight.nbytes].copy_(weight.view(-1).view(torch.uint8), non_blocking=True)
+ offset += weight.nbytes
+
+ # send the last bucket
+ torch.cuda.synchronize()
+ self.socket.send_pyobj({"bucket_meta": bucket_meta, "is_last": True})
+ self.socket.recv()
+ finally:
+ self._cleanup()
+
+ def _init_socket(self):
+ """Initialize ZMQ REQ socket and bind."""
+ self.socket = self.zmq_context.socket(zmq.REQ)
+ self.socket.bind(self.zmq_handle)
+
+ def _init_buffer(self):
+ """build communication buffer"""
+ buffer, shm = None, None
+ if not self.use_shm:
+ buffer = torch.empty(
+ self.bucket_size,
+ dtype=torch.uint8,
+ device=f"{get_current_device_name()}:{get_current_device_id()}",
+ )
+ handle = reduce_tensor(buffer)
+ self.socket.send_pyobj(handle)
+ else:
+ import uuid
+
+ # Create unique name for shared memory
+ shm_name = f"verl_weights_{uuid.uuid4().hex}"
+ shm = create_shared_memory(self.bucket_size, shm_name)
+ buffer = torch.frombuffer(shm.buf, dtype=torch.uint8)
+
+ comm_metadata = {"name": shm_name, "size": self.bucket_size}
+ self.socket.send_pyobj(comm_metadata)
+
+ self.socket.recv()
+ self.buffer = buffer
+ self.shm = shm
+
+ def _cleanup(self):
+ """clean up"""
+ if self.socket is not None:
+ self.socket.close()
+ self.socket = None
+ del self.buffer
+ self.buffer = None
+ if self.shm is not None:
+ self.shm.close()
+ self.shm.unlink()
+ del self.shm
+ self.shm = None
+ gc.collect()
+ torch.cuda.ipc_collect()
+ torch.cuda.empty_cache()
+
+
+class BucketedWeightReceiver:
+ """
+ Receive model weights via bucketed IPC transfer over ZMQ.
+
+ Receives weight tensors from BucketedWeightSender and passes each
+ bucket to a callback for processing (e.g., loading into the model).
+
+ Args:
+ zmq_handle: ZMQ IPC socket path (must match sender)
+ device: Target device for received tensors
+ use_shm: Use shared memory instead of CUDA IPC
+ """
+
+ def __init__(
+ self,
+ zmq_handle: str,
+ device: torch.device,
+ use_shm: bool = False,
+ ):
+ self.zmq_handle = zmq_handle
+ self.device = device
+ self.use_shm = use_shm
+
+ self.zmq_context = zmq.Context.instance()
+ self.socket = None
+ self.buffer = None
+ self.shm = None
+
+ def receive_weights(self, on_bucket_received: callable):
+ """
+ Receive weights from sender and process each bucket via callback.
+
+ Args:
+ on_bucket_received: Callback function(weight_dict: dict[str, torch.Tensor]) called per bucket.
+ """
+ try:
+ self._init_socket()
+ self._init_buffer()
+
+ # receive bucket and update weights
+ while True:
+ metadata = self.socket.recv_pyobj()
+ weights, tensor = [], None
+ for name, meta in metadata["bucket_meta"].items():
+ shape, dtype, offset = meta["shape"], meta["dtype"], meta["offset"]
+ size = dtype.itemsize * shape.numel()
+ # NOTE: we need to clone the tensor to release CUDA IPC memory
+ # but for shared memory, it's not necessary and if we do clone,
+ # it will cause extra memory copy overhead and slow down the process.
+ tensor = self.buffer[offset : offset + size].view(dtype=dtype).view(shape)
+ if not self.use_shm:
+ tensor = tensor.clone()
+ else:
+ tensor = tensor.to(self.device)
+ weights.append((name, tensor))
+ torch.cuda.synchronize()
+ self.socket.send(b"")
+ on_bucket_received(dict(weights))
+ del weights, tensor
+ if metadata["is_last"]:
+ break
+ finally:
+ self._cleanup()
+
+ def _init_socket(self):
+ """Initialize ZMQ REP socket and connect."""
+ self.socket = self.zmq_context.socket(zmq.REP)
+ self.socket.connect(self.zmq_handle)
+
+ def _init_buffer(self):
+ """Receive and rebuild communication buffer from sender."""
+ comm_metadata = self.socket.recv_pyobj()
+ buffer, shm = None, None
+ if not self.use_shm:
+ handle = comm_metadata
+ buffer = rebuild_cuda_ipc_tensor(handle, self.device.index)
+ assert buffer.dtype == torch.uint8
+ else:
+ shm_name = comm_metadata["name"]
+ shm_size = comm_metadata["size"]
+ buffer, shm = rebuild_shared_memory(shm_name, shm_size, dtype=torch.uint8)
+ self.socket.send(b"")
+ self.buffer = buffer
+ self.shm = shm
+
+ def _cleanup(self):
+ """clean up"""
+ if self.socket is not None:
+ self.socket.close()
+ self.socket = None
+ # Synchronize before releasing the buffer to ensure all async ops
+ # referencing it (e.g. clone, .to()) have completed.
+ torch.cuda.synchronize()
+ del self.buffer
+ self.buffer = None
+ if self.shm is not None:
+ self.shm.close()
+ del self.shm
+ self.shm = None
+ gc.collect()
+ torch.cuda.ipc_collect()
+ torch.cuda.empty_cache()
diff --git a/lightllm/utils/rl/serialization.py b/lightllm/utils/rl/serialization.py
new file mode 100644
index 0000000000..e5be3f0968
--- /dev/null
+++ b/lightllm/utils/rl/serialization.py
@@ -0,0 +1,140 @@
+"""
+Serialization helpers for RL weight-update requests.
+
+The tensor update endpoint passes CUDA tensors across processes by serializing
+their multiprocessing IPC handles rather than copying tensor data into the HTTP
+payload. This module wraps ForkingPickler for that handoff and uses a guarded
+unpickler because the payload may come from a trainer process.
+
+Copied from:
+https://github.com/sgl-project/sglang/blob/main/python/sglang/srt/utils/common.py
+"""
+import base64
+import pickle
+import io
+from dataclasses import dataclass
+from multiprocessing.reduction import ForkingPickler
+from typing import List
+
+
+class MultiprocessingSerializer:
+ @staticmethod
+ def serialize(obj, output_str: bool = False):
+ """
+ Serialize a Python object using ForkingPickler.
+
+ Args:
+ obj: The object to serialize.
+ output_str (bool): If True, return a base64-encoded string instead of raw bytes.
+
+ Returns:
+ bytes or str: The serialized object.
+ """
+ buf = io.BytesIO()
+ ForkingPickler(buf).dump(obj)
+ buf.seek(0)
+ output = buf.read()
+
+ if output_str:
+ # Convert bytes to base64-encoded string
+ output = base64.b64encode(output).decode("utf-8")
+
+ return output
+
+ @staticmethod
+ def deserialize(data):
+ """
+ Deserialize a previously serialized object.
+
+ Args:
+ data (bytes or str): The serialized data, optionally base64-encoded.
+
+ Returns:
+ The deserialized Python object.
+ """
+ if isinstance(data, str):
+ # Decode base64 string to bytes
+ data = base64.b64decode(data, validate=True)
+
+ return SafeUnpickler(io.BytesIO(data)).load()
+
+
+class SafeUnpickler(pickle.Unpickler):
+ ALLOWED_MODULE_PREFIXES = {
+ # --- Python types ---
+ "builtins.",
+ "collections.",
+ "copyreg.",
+ "functools.",
+ "itertools.",
+ "operator.",
+ "types.",
+ "weakref.",
+ # --- PyTorch types ---
+ "torch.",
+ "torch._tensor.",
+ "torch.storage.",
+ "torch.nn.parameter.",
+ "torch.autograd.function.",
+ # --- torch distributed ---
+ "torch.distributed.",
+ "torch.distributed._shard.",
+ "torch.distributed._composable.",
+ "torch._C._distributed_c10d.",
+ "torch._C._distributed_fsdp.",
+ "torch.distributed.optim.",
+ # --- multiprocessing ---
+ "multiprocessing.resource_sharer.",
+ "multiprocessing.reduction.",
+ "pickletools.",
+ # --- PEFT / LoRA ---
+ "peft.",
+ "transformers.",
+ "huggingface_hub.",
+ # --- SGLang & Unitest ---
+ "sglang.srt.weight_sync.tensor_bucket.",
+ "sglang.srt.model_executor.model_runner.",
+ "sglang.srt.layers.",
+ "sglang.srt.utils.",
+ # --- LightLLM ---
+ "lightllm.utils.",
+ }
+
+ DENY_CLASSES = {
+ ("builtins", "eval"),
+ ("builtins", "exec"),
+ ("builtins", "compile"),
+ ("os", "system"),
+ ("subprocess", "Popen"),
+ ("subprocess", "run"),
+ ("codecs", "decode"),
+ ("types", "CodeType"),
+ ("types", "FunctionType"),
+ }
+
+ def find_class(self, module, name):
+ # Block deterministic attacks
+ if (module, name) in self.DENY_CLASSES:
+ raise RuntimeError(
+ f"Blocked unsafe class loading ({module}.{name}), " f"to prevent exploitation of CVE-2025-10164"
+ )
+ # Allowlist of safe-to-load modules.
+ if any((module + ".").startswith(prefix) for prefix in self.ALLOWED_MODULE_PREFIXES):
+ return super().find_class(module, name)
+
+ # Block everything else. (Potential attack surface)
+ raise RuntimeError(
+ f"Blocked unsafe class loading ({module}.{name}), " f"to prevent exploitation of CVE-2025-10164"
+ )
+
+
+@dataclass
+class LocalSerializedTensor:
+ """torch.Tensor that gets serialized by MultiprocessingSerializer
+ (which only serializes a pointer and not the data).
+ The i-th element in the list corresponds to i-th rank's GPU."""
+
+ values: List[bytes]
+
+ def get(self, rank: int):
+ return MultiprocessingSerializer.deserialize(self.values[rank])
diff --git a/lightllm/utils/rl/tensor_bucket.py b/lightllm/utils/rl/tensor_bucket.py
new file mode 100644
index 0000000000..72defe9859
--- /dev/null
+++ b/lightllm/utils/rl/tensor_bucket.py
@@ -0,0 +1,111 @@
+"""
+Flattened tensor buckets used by RL weight-update requests.
+
+Some trainer-side integrations send many named tensors as one flattened byte
+tensor plus metadata. The server reconstructs the named tensors before passing
+them to model weight loading.
+
+Copied from:
+https://raw.githubusercontent.com/sgl-project/sglang/refs/heads/main/python/sglang/srt/weight_sync/tensor_bucket.py
+"""
+from dataclasses import dataclass
+from typing import List, Tuple
+
+import torch
+
+
+@dataclass
+class FlattenedTensorMetadata:
+ """Metadata for a tensor in a flattened bucket"""
+
+ name: str
+ shape: torch.Size
+ dtype: torch.dtype
+ start_idx: int
+ end_idx: int
+ numel: int
+
+
+class FlattenedTensorBucket:
+ """
+ A bucket that flattens multiple tensors into a single tensor for efficient processing
+ while preserving all metadata needed for reconstruction.
+ """
+
+ # This field is solely for users of to check whether the class supports this feature
+ supports_multi_dtypes = True
+
+ def __init__(
+ self,
+ named_tensors: List[Tuple[str, torch.Tensor]] = None,
+ flattened_tensor: torch.Tensor = None,
+ metadata: List[FlattenedTensorMetadata] = None,
+ ):
+ """
+ Initialize a tensor bucket from a list of named tensors OR from pre-flattened data.
+ Args:
+ named_tensors: List of (name, tensor) tuples (for creating new bucket)
+ flattened_tensor: Pre-flattened tensor (for reconstruction)
+ metadata: Pre-computed metadata (for reconstruction)
+ """
+ if named_tensors is not None:
+ # Create bucket from named tensors
+ self.metadata: List[FlattenedTensorMetadata] = [None] * len(named_tensors)
+ self.flattened_tensor: torch.Tensor = None
+
+ if not named_tensors:
+ raise ValueError("Cannot create empty tensor bucket")
+
+ # Collect metadata and flatten tensors
+ current_idx = 0
+ flattened_tensors: List[torch.Tensor] = [None] * len(named_tensors)
+
+ for i, (name, tensor) in enumerate(named_tensors):
+ flattened = tensor.flatten().view(torch.uint8)
+ flattened_tensors[i] = flattened
+
+ # Store metadata
+
+ numel = flattened.numel()
+ metadata_obj = FlattenedTensorMetadata(
+ name=name,
+ shape=tensor.shape,
+ dtype=tensor.dtype,
+ start_idx=current_idx,
+ end_idx=current_idx + numel,
+ numel=numel,
+ )
+ self.metadata[i] = metadata_obj
+ current_idx += numel
+
+ # Concatenate all flattened tensors
+ self.flattened_tensor = torch.cat(flattened_tensors, dim=0)
+ else:
+ # Initialize from pre-flattened data
+ if flattened_tensor is None or metadata is None:
+ raise ValueError("Must provide either named_tensors or both flattened_tensor and metadata")
+ self.flattened_tensor = flattened_tensor
+ self.metadata = metadata
+
+ def get_flattened_tensor(self) -> torch.Tensor:
+ """Get the flattened tensor containing all bucket tensors"""
+ return self.flattened_tensor
+
+ def get_metadata(self) -> List[FlattenedTensorMetadata]:
+ """Get metadata for all tensors in the bucket"""
+ return self.metadata
+
+ def reconstruct_tensors(self) -> List[Tuple[str, torch.Tensor]]:
+ """
+ Reconstruct original tensors from flattened tensor with optimized performance.
+ Uses memory-efficient operations to minimize allocations and copies.
+ """
+ # preallocate the result list
+ reconstructed = [None] * len(self.metadata)
+
+ for i, meta in enumerate(self.metadata):
+ tensor = self.flattened_tensor[meta.start_idx : meta.end_idx].view(meta.dtype).reshape(meta.shape)
+
+ reconstructed[i] = (meta.name, tensor)
+
+ return reconstructed
diff --git a/lightllm/utils/rl/torch_cuda_ipc.py b/lightllm/utils/rl/torch_cuda_ipc.py
new file mode 100644
index 0000000000..f2158c7a43
--- /dev/null
+++ b/lightllm/utils/rl/torch_cuda_ipc.py
@@ -0,0 +1,119 @@
+"""
+Patch torch multiprocessing CUDA tensor reductions to make cross-process
+CUDA IPC robust when different processes have different CUDA_VISIBLE_DEVICES.
+
+Torch serializes CUDA tensor IPC handles with a device index. That index is
+local to each process, so sender cuda:0 and receiver cuda:0 may refer to
+different physical GPUs. We replace the serialized device index with the GPU
+UUID on send, then map that UUID back to the receiver's local device index
+while rebuilding the tensor.
+
+The patch wraps torch's original reducers and only changes the device argument,
+so it avoids copying torch's CUDA IPC serialization implementation.
+
+Copied from:
+https://github.com/sgl-project/sglang/blob/main/python/sglang/srt/utils/patch_torch.py
+"""
+from contextlib import contextmanager
+from contextvars import ContextVar
+from typing import Callable, Union
+
+import torch
+from torch.multiprocessing import reductions
+
+
+def monkey_patch_torch_reductions():
+ """Monkey patching before Torch https://github.com/pytorch/pytorch/pull/149248 is fixed"""
+
+ # Currently, NPU does not support UUID. This has been temporarily commented out,
+ # with support expected in the fourth quarter.
+ # if _is_npu:
+ # return
+
+ if hasattr(reductions, "_reduce_tensor_original"):
+ return
+
+ reductions._reduce_tensor_original = reductions.reduce_tensor
+ reductions._rebuild_cuda_tensor_original = reductions.rebuild_cuda_tensor
+
+ reductions.reduce_tensor = _reduce_tensor_modified
+ reductions.rebuild_cuda_tensor = _rebuild_cuda_tensor_modified
+
+ reductions.init_reductions()
+
+
+# The torch CUDA IPC rebuild signature has kept the device argument at this
+# index for years. Keep this constant in one place because both the global
+# monkey patch and local bucketed IPC rebuild path need to rewrite it.
+CUDA_IPC_REBUILD_DEVICE_ARG_INDEX = 6
+_rebuild_device_fallback: ContextVar[Union[int, None]] = ContextVar("rebuild_device_fallback", default=None)
+
+
+@contextmanager
+def cuda_rebuild_device_fallback(device: Union[int, None]):
+ token = _rebuild_device_fallback.set(device)
+ try:
+ yield
+ finally:
+ _rebuild_device_fallback.reset(token)
+
+
+def _reduce_tensor_modified(*args, **kwargs):
+ output_fn, output_args = reductions._reduce_tensor_original(*args, **kwargs)
+ output_args = _modify_tuple(output_args, CUDA_IPC_REBUILD_DEVICE_ARG_INDEX, cuda_device_to_uuid)
+ return output_fn, output_args
+
+
+def _rebuild_cuda_tensor_modified(*args):
+ args = _modify_tuple(args, CUDA_IPC_REBUILD_DEVICE_ARG_INDEX, cuda_device_from_maybe_uuid)
+ return reductions._rebuild_cuda_tensor_original(*args)
+
+
+def get_current_device_name() -> str:
+ if torch.cuda.is_available():
+ return "cuda"
+ return "cpu"
+
+
+def get_current_device_module():
+ device_name = get_current_device_name()
+ try:
+ return getattr(torch, device_name)
+ except AttributeError:
+ return torch.cuda
+
+
+def get_current_device_id() -> int:
+ return get_current_device_module().current_device()
+
+
+def cuda_device_to_uuid(device: int) -> str:
+ return str(torch.cuda.get_device_properties(device).uuid)
+
+
+def cuda_device_from_maybe_uuid(device_maybe_uuid: Union[int, str]) -> int:
+ if isinstance(device_maybe_uuid, int):
+ return device_maybe_uuid
+
+ if isinstance(device_maybe_uuid, str):
+ for device in range(torch.cuda.device_count()):
+ if str(torch.cuda.get_device_properties(device).uuid) == device_maybe_uuid:
+ return device
+ fallback_device = _rebuild_device_fallback.get()
+ if fallback_device is not None:
+ return fallback_device
+ raise Exception("Invalid device_uuid=" + device_maybe_uuid)
+
+ raise Exception(f"Unknown type: {device_maybe_uuid=}")
+
+
+def rebuild_cuda_ipc_tensor(handle: tuple[Callable, tuple], device_id: Union[int, None] = None) -> torch.Tensor:
+ func, args = handle
+ list_args = list(args)
+ if device_id is not None:
+ list_args[CUDA_IPC_REBUILD_DEVICE_ARG_INDEX] = device_id
+ return func(*list_args)
+
+
+def _modify_tuple(t, index: int, modifier: Callable):
+ return *t[:index], modifier(t[index]), *t[index + 1 :]
diff --git a/lightllm/utils/shm_port_args.py b/lightllm/utils/shm_port_args.py
new file mode 100644
index 0000000000..44f56fcc2d
--- /dev/null
+++ b/lightllm/utils/shm_port_args.py
@@ -0,0 +1,324 @@
+"""
+ShmPortArgs: 跨进程共享的启动端口表(Shared Memory Port Args)。
+
+设计目的
+--------
+LightLLM 启动时有两类端口:
+ [user_set] 用户通过 CLI / 默认值已经写在 start args 里(如 --port、--pd_master_port)。
+ 未设置时访问会直接报错,不会动态分配。
+ [dynamic] args 中为 None,需要运行时动态申请(如 router_port、metric_port)。
+
+本类把这两类端口统一成属性访问接口,并在进程间通过命名 POSIX shm 共享
+「已动态分配」的结果,避免启动阶段预先批量占坑,也避免多进程重复分配冲突。
+
+存储与并发
+----------
+- shm 名:`{unique_server_name}_shm_port_args`
+- 内容:pickle 序列化的 dict[str, int | list[int]]
+- 分配:socket.bind("", 0) 取空闲端口;分配前排除
+ 1) start args 中已设置的 port 字段
+ 2) 本 shm 中已经分配过的端口
+- 互斥:FileLock(`/tmp/{shm_name}.lock`)
+
+使用约定
+--------
+使用 ShmPortArgs 之前,必须先初始化以下环境信息(否则会直接抛错):
+ 1. `set_unique_server_name(args)`
+ → 写入 `LIGHTLLM_UNIQUE_SERVICE_NAME_ID`,供 `get_unique_server_name()` 使用,用于拼 shm 名。
+ 2. `set_env_start_args(args)`
+ → 写入 `LIGHTLLM_START_ARGS`,供 `get_env_start_args()` 使用,用于读取用户已设置端口、
+ 以及 visual_dp / audio_dp 等分配参数。
+
+之后:
+ 3. 启动主进程:`ShmPortArgs.get_instance(create=True)` 创建空表。
+ 4. 同进程其它位置 / 子进程:`ShmPortArgs.get_instance()` 复用或 link 同一块 shm。
+ 5. 请通过 `get_instance` / `get_shm_port_args` 获取对象;直接构造会绕过进程内单例。
+
+示例
+----
+ set_unique_server_name(args) # 初始化 get_unique_server_name 所需环境变量
+ set_env_start_args(args) # 初始化 get_env_start_args 所需环境变量
+
+ ShmPortArgs.get_instance(create=True)
+ ports = get_shm_port_args()
+ router_port = ports.router_port # dynamic:按需分配
+ http_port = ports.port # user_set:读 args
+ vit_ports = ports.visual_nccl_ports # dynamic list
+"""
+
+from __future__ import annotations
+
+import os
+import pickle
+import socket
+import struct
+from typing import Dict, List, Set, Union
+
+from filelock import FileLock
+
+from lightllm.utils.envs_utils import get_env_start_args, get_unique_server_name
+from lightllm.utils.log_utils import init_logger
+from lightllm.utils.shm_utils import create_or_link_shm
+
+logger = init_logger(__name__)
+
+PortValue = Union[int, List[int]]
+
+
+class ShmPortArgs:
+ _SHM_SIZE = 64 * 1024
+ _ALLOC_MAX_RETRY = 128
+ _instance: "ShmPortArgs | None" = None
+
+ def __init__(self, create: bool = False):
+ uni = get_unique_server_name()
+ if not uni:
+ raise RuntimeError(
+ "LIGHTLLM_UNIQUE_SERVICE_NAME_ID is unset; " "call set_unique_server_name(args) before ShmPortArgs"
+ )
+ if "LIGHTLLM_START_ARGS" not in os.environ:
+ raise RuntimeError("LIGHTLLM_START_ARGS is unset; call set_env_start_args(args) before ShmPortArgs")
+
+ self._shm_name = f"{uni}_shm_port_args"
+ self._lock = FileLock(f"/tmp/{self._shm_name}.lock")
+ self.shm = create_or_link_shm(
+ self._shm_name,
+ self._SHM_SIZE,
+ force_mode="create" if create else "link",
+ auto_cleanup=create,
+ )
+ if create:
+ self._save({})
+
+ @classmethod
+ def get_instance(cls, create: bool = False) -> "ShmPortArgs":
+ """同进程单例;create 仅在首次构造时生效。"""
+ if cls._instance is None:
+ cls._instance = cls(create=create)
+ return cls._instance
+
+ # =====================================================================
+ # [user_set] 用户已在 args 中设置(含 CLI 默认值);未设置则报错
+ # =====================================================================
+
+ # [user_set] HTTP API listen port
+ @property
+ def port(self) -> int:
+ return self._get_user_set_port("port")
+
+ # [user_set] PD master port
+ @property
+ def pd_master_port(self) -> int:
+ return self._get_user_set_port("pd_master_port")
+
+ # [user_set] config server port
+ @property
+ def config_server_port(self) -> int:
+ return self._get_user_set_port("config_server_port")
+
+ # [user_set] config server visual redis port
+ @property
+ def config_server_visual_redis_port(self) -> int:
+ return self._get_user_set_port("config_server_visual_redis_port")
+
+ # [user_set] multinode http manager port
+ @property
+ def multinode_httpmanager_port(self) -> int:
+ return self._get_user_set_port("multinode_httpmanager_port")
+
+ # [user_set] multinode router gloo port
+ @property
+ def multinode_router_gloo_port(self) -> int:
+ return self._get_user_set_port("multinode_router_gloo_port")
+
+ # =====================================================================
+ # [dynamic] args 可为 None,需要时动态分配并写入 shm
+ # =====================================================================
+
+ # [dynamic] pytorch distributed / NCCL TCPStore port
+ @property
+ def nccl_port(self) -> int:
+ return self._get_from_args_or_alloc("nccl_port")
+
+ # [dynamic] visual-only RPyC port
+ @property
+ def visual_rpyc_port(self) -> int:
+ return self._get_from_args_or_alloc("visual_rpyc_port")
+
+ # [dynamic] visual vit nccl ports (list, len=visual_dp)
+ @property
+ def visual_nccl_ports(self) -> List[int]:
+ return self._get_from_args_or_alloc("visual_nccl_ports", count=int(get_env_start_args().visual_dp))
+
+ # [dynamic] audio encoder nccl ports (list, len=audio_dp)
+ @property
+ def audio_nccl_ports(self) -> List[int]:
+ return self._get_from_args_or_alloc("audio_nccl_ports", count=int(get_env_start_args().audio_dp))
+
+ # [dynamic] router zmq port
+ @property
+ def router_port(self) -> int:
+ return self._get_from_args_or_alloc("router_port")
+
+ # [dynamic] router profiler zmq port
+ @property
+ def router_profiler_port(self) -> int:
+ return self._get_from_args_or_alloc("router_profiler_port")
+
+ # [dynamic] detokenization zmq port
+ @property
+ def detokenization_port(self) -> int:
+ return self._get_from_args_or_alloc("detokenization_port")
+
+ # [dynamic] http server internal zmq port
+ @property
+ def http_server_port(self) -> int:
+ return self._get_from_args_or_alloc("http_server_port")
+
+ # [dynamic] visual server zmq port
+ @property
+ def visual_port(self) -> int:
+ return self._get_from_args_or_alloc("visual_port")
+
+ # [dynamic] audio server zmq port
+ @property
+ def audio_port(self) -> int:
+ return self._get_from_args_or_alloc("audio_port")
+
+ # [dynamic] embed cache rpyc port
+ @property
+ def cache_port(self) -> int:
+ return self._get_from_args_or_alloc("cache_port")
+
+ # [dynamic] metrics rpyc port
+ @property
+ def metric_port(self) -> int:
+ return self._get_from_args_or_alloc("metric_port")
+
+ # [dynamic] multi-level kv cache port
+ @property
+ def multi_level_kv_cache_port(self) -> int:
+ return self._get_from_args_or_alloc("multi_level_kv_cache_port")
+
+ # [dynamic] router RL RPyC port
+ @property
+ def rl_rpyc_port(self) -> int:
+ return self._get_from_args_or_alloc("rl_rpyc_port")
+
+ def close(self) -> None:
+ if self.shm is not None:
+ self.shm.close()
+ self.shm = None
+ if ShmPortArgs._instance is self:
+ ShmPortArgs._instance = None
+
+ def _get_user_set_port(self, name: str) -> int:
+ """只读 args;未设置则报错,不动态分配。"""
+ value = getattr(get_env_start_args(), name, None)
+ if value is None:
+ raise RuntimeError(f"user_set port '{name}' is None; set it in start args before use")
+ return int(value)
+
+ def _get_from_args_or_alloc(self, name: str, count: int = 1) -> PortValue:
+ """args 已设置则用用户值,否则写入同一份 shm 动态分配。"""
+ value = getattr(get_env_start_args(), name, None)
+ if value is not None:
+ if count == 1:
+ return int(value)
+ return [int(v) for v in value[:count]]
+ return self._get_or_alloc(name, count=count)
+
+ def _get_or_alloc(self, name: str, count: int = 1) -> PortValue:
+ with self._lock:
+ ports = self._load()
+ if name in ports:
+ return ports[name]
+
+ # 分配时必须同时排除:
+ # 1) args 中用户已设置的端口
+ # 2) 本 shm 中已经分配过的端口
+ reserved = self._ports_from_args() | self._ports_from_shm(ports)
+
+ if count == 1:
+ port = self._alloc_free_port(reserved)
+ ports[name] = port
+ self._save(ports)
+ logger.info(f"ShmPortArgs alloc {name}={port}")
+ return port
+
+ allocated: List[int] = []
+ for _ in range(count):
+ port = self._alloc_free_port(reserved)
+ reserved.add(port) # 本轮后续分配也要排除刚分到的端口
+ allocated.append(port)
+ ports[name] = allocated
+ self._save(ports)
+ logger.info(f"ShmPortArgs alloc {name}={allocated}")
+ return allocated
+
+ @staticmethod
+ def _ports_from_args() -> Set[int]:
+ """遍历 args,收集 key 含 port 且 value 非 None 的端口,分配时必须避开。"""
+ args = get_env_start_args()
+ reserved: Set[int] = set()
+ for key, value in dict(args).items():
+ if "port" not in key.lower() or value is None:
+ continue
+ if isinstance(value, (list, tuple)):
+ reserved.update(int(v) for v in value if v is not None)
+ else:
+ reserved.add(int(value))
+ return reserved
+
+ @staticmethod
+ def _ports_from_shm(ports: Dict[str, PortValue]) -> Set[int]:
+ """收集 shm 中已经分配过的端口。"""
+ reserved: Set[int] = set()
+ for value in ports.values():
+ if isinstance(value, list):
+ reserved.update(int(v) for v in value)
+ else:
+ reserved.add(int(value))
+ return reserved
+
+ def _load(self) -> Dict[str, PortValue]:
+ n = struct.unpack_from("I", self.shm.buf, 0)[0]
+ if n == 0:
+ return {}
+ return pickle.loads(bytes(self.shm.buf[4 : 4 + n]))
+
+ def _save(self, ports: Dict[str, PortValue]) -> None:
+ blob = pickle.dumps(ports, protocol=pickle.HIGHEST_PROTOCOL)
+ if 4 + len(blob) > self._SHM_SIZE:
+ raise RuntimeError(f"port table too large: {len(blob)} bytes")
+ struct.pack_into("I", self.shm.buf, 0, len(blob))
+ self.shm.buf[4 : 4 + len(blob)] = blob
+
+ def _alloc_free_port(self, reserved: Set[int]) -> int:
+ for _ in range(self._ALLOC_MAX_RETRY):
+ with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as sock:
+ sock.bind(("", 0))
+ port = int(sock.getsockname()[1])
+ if port in reserved:
+ logger.warning(f"skip port {port}, already in args or shm-allocated")
+ continue
+ if not self._is_port_free(port):
+ logger.warning(f"skip port {port}, not free")
+ continue
+ return port
+ raise RuntimeError(f"failed to allocate free port after {self._ALLOC_MAX_RETRY} retries, reserved={reserved}")
+
+ @staticmethod
+ def _is_port_free(port: int) -> bool:
+ with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as sock:
+ sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
+ try:
+ sock.bind(("", port))
+ return True
+ except OSError:
+ return False
+
+
+def get_shm_port_args(create: bool = False) -> ShmPortArgs:
+ """Convenience accessor for the process-local ShmPortArgs singleton."""
+ return ShmPortArgs.get_instance(create=create)
diff --git a/lightllm/utils/torch_memory_saver_utils.py b/lightllm/utils/torch_memory_saver_utils.py
new file mode 100644
index 0000000000..e0e6cad94a
--- /dev/null
+++ b/lightllm/utils/torch_memory_saver_utils.py
@@ -0,0 +1,105 @@
+import torch
+from contextlib import contextmanager
+from enum import Enum
+from lightllm.utils.log_utils import init_logger
+
+try:
+ from torch_memory_saver import (
+ torch_memory_saver,
+ configure_subprocess as tms_configure_subprocess,
+ )
+
+ HAS_TORCH_MEMORY_SAVER = True
+
+except ImportError:
+ HAS_TORCH_MEMORY_SAVER = False
+ pass
+
+logger = init_logger(__name__)
+
+
+class MemoryTag(Enum):
+ # torch_memory_saver 通过 tag 区分不同类型的显存区域,后续 pause/resume
+ # 可以只针对某一类内存做释放和恢复。
+ KV_CACHE = "kv_cache"
+ WEIGHT = "weights"
+ GRAPH = "graph"
+
+ def is_kv_cache(self):
+ return self == MemoryTag.KV_CACHE
+
+ def is_weight(self):
+ return self == MemoryTag.WEIGHT
+
+ def is_graph(self):
+ return self == MemoryTag.GRAPH
+
+ def __str__(self):
+ return self.value
+
+
+class TorchMemorySaverWrapper:
+ # 统一返回真实实现或空实现,调用方不需要到处判断
+ # enable_torch_memory_saver 是否开启。
+ def __new__(cls, enable_torch_memory_saver: bool = False):
+ if enable_torch_memory_saver:
+ assert (
+ HAS_TORCH_MEMORY_SAVER
+ ), "torch_memory_saver is not installed, please install it via `pip install torch_memory_saver`."
+ return _TorchMemorySaver()
+ else:
+ return _TorchMemorySaverFake()
+
+
+class _TorchMemorySaver:
+ @contextmanager
+ def configure_subprocess(self):
+ # 子进程启动需要放在该上下文里,让 torch_memory_saver 在 worker
+ # 进程中完成必要的初始化。
+ with tms_configure_subprocess():
+ yield
+
+ def region(self, tag: MemoryTag, enable_cpu_backup: bool = False):
+ # 记录这个上下文内产生的显存分配;enable_cpu_backup 用于需要
+ # pause 后还能恢复内容的区域,比如权重。
+ return torch_memory_saver.region(tag=tag.value, enable_cpu_backup=enable_cpu_backup)
+
+ def cuda_graph(self, graph_obj: torch.cuda.CUDAGraph, **kwargs):
+ # CUDA graph 的 private pool 也单独打 tag,避免和普通权重/KV cache
+ # 的显存管理混在一起。
+ return torch_memory_saver.cuda_graph(cuda_graph=graph_obj, **kwargs, tag=MemoryTag.GRAPH.value)
+
+ def disable(self):
+ return torch_memory_saver.disable()
+
+ def pause(self, tag: MemoryTag):
+ return torch_memory_saver.pause(tag=tag.value)
+
+ def resume(self, tag: MemoryTag):
+ return torch_memory_saver.resume(tag=tag.value)
+
+
+class _TorchMemorySaverFake:
+ # 未开启 torch_memory_saver 时保持相同接口,保证调用方逻辑完全一致。
+ @contextmanager
+ def configure_subprocess(self):
+ yield
+
+ @contextmanager
+ def region(self, tag: MemoryTag, enable_cpu_backup: bool = False):
+ yield
+
+ def cuda_graph(self, graph_obj: torch.cuda.CUDAGraph, **kwargs):
+ return torch.cuda.graph(graph_obj, **kwargs)
+
+ @contextmanager
+ def disable(self):
+ yield
+
+ def pause(self, tag: MemoryTag):
+ logger.warning("torch_memory_saver is not enabled, pause is not supported.")
+ return
+
+ def resume(self, tag: MemoryTag):
+ logger.warning("torch_memory_saver is not enabled, resume is not supported.")
+ return
diff --git a/requirements.txt b/requirements.txt
index 603d0e488f..180f9d3ced 100644
--- a/requirements.txt
+++ b/requirements.txt
@@ -98,3 +98,4 @@ nixl==1.2.0
xformers==0.0.35
redis==7.3.0
litellm>=1.52.0,<1.85
+torch_memory_saver==0.0.9.post1
\ No newline at end of file
diff --git a/test/test_api/test_abort_chaos.py b/test/test_api/test_abort_chaos.py
new file mode 100644
index 0000000000..63f226717e
--- /dev/null
+++ b/test/test_api/test_abort_chaos.py
@@ -0,0 +1,105 @@
+"""
+Two-stage abort test against a running lightllm server.
+
+Stage 1: spawn N concurrent streams, then post /abort_request abort_all=True;
+ verify every stream terminates quickly.
+Stage 2: spawn N concurrent streams; each stream is independently assigned a
+ random fate (disconnect mid-stream or run to completion). The server
+ must keep serving the survivors and stay healthy afterwards.
+
+Usage:
+ python test/test_api/test_abort_chaos.py --url http://127.0.0.1:8000
+"""
+
+import argparse
+import asyncio
+import json
+import random
+import time
+from collections import Counter
+
+import httpx
+
+
+PROMPTS = [
+ "Write a long detailed essay about the history of computing.",
+ "Tell me a long story about dragons and knights.",
+ "Explain quantum mechanics in detail with lots of examples.",
+ "Describe the plot of a 5-part fantasy novel series.",
+ "Compose a long poem about the seasons.",
+]
+
+
+async def stream_task(client, url, mode, max_new_tokens):
+ payload = {
+ "inputs": random.choice(PROMPTS),
+ "parameters": {"max_new_tokens": max_new_tokens, "temperature": 0.7, "do_sample": True},
+ }
+ drop_after = random.randint(20, 200)
+ tokens = 0
+ finish_reason = None
+ t0 = time.time()
+ try:
+ async with client.stream("POST", f"{url}/generate_stream", json=payload, timeout=180.0) as r:
+ async for line in r.aiter_lines():
+ tokens += 1
+ if line.startswith("data:"):
+ chunk = json.loads(line[len("data:") :])
+ if chunk.get("finished"):
+ finish_reason = chunk.get("finish_reason")
+ if mode == "disconnect" and tokens >= drop_after:
+ break
+ return (mode, finish_reason or "ok", tokens, time.time() - t0)
+ except Exception as e:
+ return (mode, f"exc:{type(e).__name__}", tokens, time.time() - t0)
+
+
+def summarize(results):
+ outcomes = Counter()
+ for r in results:
+ outcomes[(r[0], r[1]) if isinstance(r, tuple) else f"raised:{type(r).__name__}"] += 1
+ for k, v in sorted(outcomes.items(), key=str):
+ print(f" {k}: {v}")
+
+
+async def stage_abort_all(client, url, concurrency, max_new_tokens):
+ print("\n===== STAGE 1: abort_all on N concurrent streams =====")
+ tasks = [asyncio.create_task(stream_task(client, url, "finish", max_new_tokens)) for _ in range(concurrency)]
+ await asyncio.sleep(2.0)
+ t0 = time.time()
+ r = await client.post(f"{url}/abort_request", json={"abort_all": True}, timeout=10.0)
+ print(f"abort_all status={r.status_code}")
+ results = await asyncio.gather(*tasks, return_exceptions=True)
+ print(f"all streams settled in {time.time() - t0:.2f}s")
+ summarize(results)
+
+
+async def stage_random_chaos(client, url, concurrency, max_new_tokens):
+ print("\n===== STAGE 2: random per-stream chaos =====")
+ modes = random.choices(["disconnect", "finish"], weights=[80, 20], k=concurrency)
+ tasks = [asyncio.create_task(stream_task(client, url, m, max_new_tokens)) for m in modes]
+ results = await asyncio.gather(*tasks, return_exceptions=True)
+ summarize(results)
+
+
+async def run(url, concurrency, max_new_tokens):
+ async with httpx.AsyncClient() as client:
+ await stage_abort_all(client, url, concurrency, max_new_tokens)
+ await stage_random_chaos(client, url, concurrency, max_new_tokens)
+ print("\nALL CHAOS TESTS PASSED")
+
+
+def main():
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--url", default="http://127.0.0.1:8000")
+ parser.add_argument("--concurrency", type=int, default=24)
+ parser.add_argument("--max_new_tokens", type=int, default=2048)
+ parser.add_argument("--seed", type=int, default=42)
+ args = parser.parse_args()
+
+ random.seed(args.seed)
+ asyncio.run(run(args.url, args.concurrency, args.max_new_tokens))
+
+
+if __name__ == "__main__":
+ main()
diff --git a/test/test_api/test_abort_request.py b/test/test_api/test_abort_request.py
new file mode 100644
index 0000000000..ca99f6298c
--- /dev/null
+++ b/test/test_api/test_abort_request.py
@@ -0,0 +1,432 @@
+"""
+Test the /abort_request endpoint against a running lightllm server.
+
+What this test asserts (and why it does not assert "stream becomes finish_reason='abort'"):
+
+ In normal / chunked_prefill mode, /abort_request:
+ - sets shm_req.is_aborted = True
+ - drives the router to send AbortedReqCmd, which sets InferReq.infer_aborted = True
+ on the worker
+ - causes still-waiting (not yet scheduled) reqs to be freed with FINISHED_ABORTED
+ - but does NOT cause already-running reqs to early-exit; they finish at max_new_tokens
+ / EOS / stop sequence as usual. (The shm flag is consumed by audio/visual servers
+ and pd_nixl mode, but the LLM inference loop never short-circuits on it.)
+
+So the test verifies the contract that actually exists today:
+
+ Stage A: bogus request_id -> HTTP 200, server log "not exist" warning
+ Stage B: abort_all on an idle server -> HTTP 200, no errors
+ Stage C: abort_all on a running stream
+ -> HTTP 200; server log shows "aborted group_request_id N" warning
+ -> the stream terminates within reasonable time (whether via abort or
+ natural max_new_tokens completion)
+ Stage D: abort by SPECIFIC request_id on a running stream
+ -> resolve the lightllm_req_id from the server log (via X-Request-Id),
+ POST /abort_request with that exact id, verify the targeted log
+ warning lands and the stream terminates
+ Stage E: server remains healthy and answers a fresh /generate
+
+Usage:
+ python test/test_api/test_abort_request.py \
+ --url http://127.0.0.1:8000 \
+ --server_log_path /tmp/lightllm_test/server.log
+"""
+
+import argparse
+import json
+import os
+import re
+import sys
+import threading
+import time
+import uuid
+from typing import List, Optional, Tuple
+
+import requests
+
+
+GREEN = "\033[32m"
+RED = "\033[31m"
+YELLOW = "\033[33m"
+RESET = "\033[0m"
+
+
+def banner(msg: str):
+ print(f"\n{YELLOW}=== {msg} ==={RESET}", flush=True)
+
+
+def ok(msg: str):
+ print(f" {GREEN}OK{RESET} {msg}", flush=True)
+
+
+def fail(msg: str):
+ print(f" {RED}FAIL{RESET} {msg}", flush=True)
+
+
+# ---------------- HTTP helpers ----------------
+
+
+def _get_health(url: str, timeout=5):
+ return requests.get(url + "/health", timeout=timeout)
+
+
+def post_abort(url: str, request_id: Optional[int] = None, abort_all: bool = False) -> Tuple[int, str]:
+ payload = {"abort_all": abort_all}
+ if request_id is not None:
+ payload["request_id"] = request_id
+ r = requests.post(url + "/abort_request", json=payload, timeout=30)
+ return r.status_code, r.text
+
+
+# ---------------- streaming helpers ----------------
+
+
+def _stream_run(
+ url: str,
+ prompt: str,
+ max_new_tokens: int,
+ x_request_id: str,
+ out: dict,
+ close_after_n: Optional[int] = None,
+):
+ """
+ Issue a /generate_stream and append every event to out["events"].
+ If close_after_n is set, the underlying socket is forcibly closed
+ (TCP RST via SO_LINGER + close) after that many events arrive — kept
+ here for completeness even though no current stage uses it. Sets
+ out["error"] on transport errors.
+ """
+ headers = {"X-Request-Id": x_request_id, "Content-Type": "application/json"}
+ body = {
+ "inputs": prompt,
+ "parameters": {
+ "max_new_tokens": max_new_tokens,
+ "do_sample": False,
+ "ignore_eos": True,
+ },
+ }
+ out["events"] = []
+ out["start"] = time.time()
+ out["error"] = None
+ out["closed_intentionally"] = False
+ try:
+ # urllib3 keeps the socket pooled; we need direct access to force-close.
+ with requests.post(url + "/generate_stream", json=body, headers=headers, stream=True, timeout=120) as r:
+ r.raise_for_status()
+ for raw in r.iter_lines(decode_unicode=True):
+ if not raw:
+ continue
+ if raw.startswith("data:"):
+ raw = raw[len("data:") :]
+ try:
+ ev = json.loads(raw)
+ except Exception:
+ continue
+ ev["_t"] = time.time() - out["start"]
+ out["events"].append(ev)
+ if close_after_n is not None and len(out["events"]) >= close_after_n:
+ out["closed_intentionally"] = True
+ # Reach into urllib3 to force a TCP RST so the server sees
+ # the disconnect immediately rather than after a graceful
+ # FIN that hypercorn might not propagate while the response
+ # is mid-stream.
+ try:
+ import socket as _socket
+
+ sock = r.raw._fp.fp.raw._sock # type: ignore[attr-defined]
+ # SO_LINGER with timeout 0 -> RST on close.
+ l_onoff, l_linger = 1, 0
+ sock.setsockopt(
+ _socket.SOL_SOCKET,
+ _socket.SO_LINGER,
+ int.to_bytes(l_onoff, 4, "little") + int.to_bytes(l_linger, 4, "little"),
+ )
+ sock.close()
+ except Exception as e:
+ out["close_error"] = repr(e)
+ break
+ if ev.get("finished"):
+ break
+ except Exception as e:
+ out["error"] = repr(e)
+ out["end"] = time.time()
+
+
+def start_stream(
+ url: str, prompt: str, max_new_tokens: int, close_after_n: Optional[int] = None
+) -> Tuple[threading.Thread, dict, str]:
+ xid = uuid.uuid4().hex
+ out = {}
+ th = threading.Thread(target=_stream_run, args=(url, prompt, max_new_tokens, xid, out, close_after_n))
+ th.daemon = True
+ th.start()
+ return th, out, xid
+
+
+def wait_for_first_token(out: dict, timeout: float = 30.0) -> bool:
+ deadline = time.time() + timeout
+ while time.time() < deadline:
+ if out.get("events"):
+ return True
+ time.sleep(0.05)
+ return False
+
+
+def get_finish_reason(out: dict) -> Optional[str]:
+ for ev in reversed(out.get("events") or []):
+ fr = ev.get("finish_reason")
+ if fr:
+ return fr
+ return None
+
+
+# ---------------- log helpers ----------------
+
+
+def _read_log_tail(server_log_path: Optional[str], max_bytes: int = 256 * 1024) -> str:
+ if not server_log_path or not os.path.exists(server_log_path):
+ return ""
+ try:
+ size = os.path.getsize(server_log_path)
+ with open(server_log_path, "rb") as f:
+ if size > max_bytes:
+ f.seek(size - max_bytes)
+ return f.read().decode("utf-8", errors="ignore")
+ except FileNotFoundError:
+ return ""
+
+
+def grep_log_for_pattern(server_log_path: Optional[str], pattern: re.Pattern, timeout: float = 5.0) -> Optional[str]:
+ """Poll the tail of the server log for a regex match."""
+ deadline = time.time() + timeout
+ while time.time() < deadline:
+ tail = _read_log_tail(server_log_path)
+ m = pattern.search(tail)
+ if m:
+ return m.group(0)
+ time.sleep(0.1)
+ return None
+
+
+def grep_log_after_offset(
+ server_log_path: Optional[str], start_offset: int, pattern: re.Pattern, timeout: float = 5.0
+) -> Optional[str]:
+ """Poll the server log starting at start_offset for a regex match.
+ Only content written after start_offset is considered, so this isolates
+ a stage from log produced by earlier stages."""
+ if not server_log_path:
+ return None
+ deadline = time.time() + timeout
+ while time.time() < deadline:
+ try:
+ with open(server_log_path, "rb") as f:
+ f.seek(start_offset)
+ new = f.read().decode("utf-8", errors="ignore")
+ except FileNotFoundError:
+ new = ""
+ m = pattern.search(new)
+ if m:
+ return m.group(0)
+ time.sleep(0.1)
+ return None
+
+
+def server_log_size(server_log_path: Optional[str]) -> int:
+ if not server_log_path or not os.path.exists(server_log_path):
+ return 0
+ return os.path.getsize(server_log_path)
+
+
+def lookup_lightllm_req_id_from_log(server_log_path: str, x_request_id: str, timeout: float = 5.0) -> Optional[int]:
+ pattern = re.compile(rf"received req X-Request-Id:{re.escape(x_request_id)}\b.*?lightllm_req_id:(\d+)")
+ deadline = time.time() + timeout
+ while time.time() < deadline:
+ tail = _read_log_tail(server_log_path)
+ m = pattern.search(tail)
+ if m:
+ return int(m.group(1))
+ time.sleep(0.1)
+ return None
+
+
+# ---------------- stages ----------------
+
+
+def stage_a_bogus_id(url: str) -> bool:
+ banner("Stage A: abort with a non-existent id")
+ bogus = 99_999_999
+ code, text = post_abort(url, request_id=bogus, abort_all=False)
+ print(f" /abort_request request_id={bogus} -> HTTP {code} body={text!r}")
+ if code != 200:
+ fail(f"expected HTTP 200, got {code}")
+ return False
+ ok("HTTP 200")
+ return True
+
+
+def stage_b_abort_all_idle(url: str) -> bool:
+ banner("Stage B: abort_all on an idle server")
+ code, text = post_abort(url, abort_all=True)
+ print(f" /abort_request abort_all=true -> HTTP {code} body={text!r}")
+ if code != 200:
+ fail(f"expected HTTP 200, got {code}")
+ return False
+ ok("HTTP 200")
+ return True
+
+
+def stage_c_abort_running(url: str, server_log_path: Optional[str]) -> bool:
+ banner("Stage C: abort_all on a running stream")
+ log_offset = server_log_size(server_log_path)
+ th, out, xid = start_stream(url, "Recite the alphabet repeatedly.", max_new_tokens=200)
+ if not wait_for_first_token(out, timeout=30.0):
+ fail("did not receive any tokens before abort")
+ return False
+ first_t = out["events"][0]["_t"]
+ ok(f"first token at +{first_t:.2f}s")
+
+ target_id = lookup_lightllm_req_id_from_log(server_log_path, xid, timeout=5.0) if server_log_path else None
+ print(f" resolved lightllm_req_id from log: {target_id}")
+
+ code, text = post_abort(url, abort_all=True)
+ print(f" /abort_request abort_all=true -> HTTP {code} body={text!r}")
+ if code != 200:
+ fail(f"expected HTTP 200, got {code}")
+ return False
+
+ th.join(timeout=60.0)
+ if th.is_alive():
+ fail("stream did not terminate within 60s of abort")
+ return False
+ fr = get_finish_reason(out)
+ n = len(out.get("events") or [])
+ print(f" stream events received: {n}, finish_reason={fr!r}, error={out.get('error')!r}")
+
+ # The api itself succeeded; whether the stream got a clean 'abort' finish reason
+ # depends on which mode-backend the server is running. We DO assert the abort
+ # warning landed in the server log though, scoped to log content produced after
+ # this stage started so we don't match earlier-stage residue.
+ if server_log_path:
+ if target_id is not None:
+ pat = re.compile(rf"aborted group_request_id {target_id}\b")
+ else:
+ pat = re.compile(r"aborted group_request_id \d+")
+ hit = grep_log_after_offset(server_log_path, log_offset, pat, timeout=5.0)
+ if not hit:
+ fail("could not find 'aborted group_request_id' in server log (post-stage)")
+ return False
+ ok(f"server log recorded: {hit!r}")
+ else:
+ print(" no --server_log_path; skipped log assertion")
+ ok("stream terminated and abort acknowledged")
+ return True
+
+
+def stage_d_abort_by_id(url: str, server_log_path: Optional[str]) -> bool:
+ banner("Stage D: abort by specific request_id on a running stream")
+ if not server_log_path:
+ print(" --server_log_path not provided; skipping (we need the log to resolve req_id)")
+ return True
+
+ log_offset = server_log_size(server_log_path)
+ th, out, xid = start_stream(url, "Sing a long lullaby for the moon.", max_new_tokens=300)
+ if not wait_for_first_token(out, timeout=30.0):
+ fail("did not receive any tokens before abort")
+ return False
+ ok(f"first token at +{out['events'][0]['_t']:.2f}s, X-Request-Id={xid[:8]}…")
+
+ target_id = lookup_lightllm_req_id_from_log(server_log_path, xid, timeout=5.0)
+ if target_id is None:
+ fail("could not resolve lightllm_req_id from server log; cannot test by-id abort")
+ return False
+ print(f" resolved lightllm_req_id: {target_id}")
+
+ code, text = post_abort(url, request_id=target_id, abort_all=False)
+ print(f" /abort_request request_id={target_id} -> HTTP {code} body={text!r}")
+ if code != 200:
+ fail(f"expected HTTP 200, got {code}")
+ return False
+
+ th.join(timeout=60.0)
+ if th.is_alive():
+ fail("stream did not terminate within 60s")
+ return False
+ fr = get_finish_reason(out)
+ n = len(out.get("events") or [])
+ print(f" stream events received: {n}, finish_reason={fr!r}")
+
+ pat = re.compile(rf"aborted group_request_id {target_id}\b")
+ hit = grep_log_after_offset(server_log_path, log_offset, pat, timeout=5.0)
+ if not hit:
+ fail(f"could not find 'aborted group_request_id {target_id}' in server log (post-stage)")
+ return False
+ ok(f"server log recorded: {hit!r}")
+ return True
+
+
+def stage_e_health_after(url: str) -> bool:
+ banner("Stage E: server still serves a normal /generate")
+ r = requests.post(
+ url + "/generate",
+ json={
+ "inputs": "The capital of France is",
+ "parameters": {"max_new_tokens": 6, "do_sample": False},
+ },
+ timeout=60,
+ )
+ print(f" /generate -> HTTP {r.status_code} {r.text[:200]}")
+ if r.status_code != 200:
+ fail(f"final /generate failed with {r.status_code}")
+ return False
+ body = r.json()
+ text = body.get("generated_text")
+ if isinstance(text, list):
+ text = text[0]
+ if not text or not text.strip():
+ fail("final /generate returned empty text")
+ return False
+ ok(f"final /generate returned {text!r}")
+ return True
+
+
+# ---------------- main ----------------
+
+
+def main():
+ ap = argparse.ArgumentParser()
+ ap.add_argument("--url", default="http://127.0.0.1:8000")
+ ap.add_argument(
+ "--server_log_path",
+ default=None,
+ help="optional path to the server stdout/stderr log; enables log-grep assertions",
+ )
+ args = ap.parse_args()
+
+ try:
+ r = _get_health(args.url)
+ r.raise_for_status()
+ except Exception as e:
+ fail(f"server at {args.url} not reachable: {e}")
+ sys.exit(1)
+ ok(f"server reachable at {args.url}")
+
+ results = []
+ results.append(("A", stage_a_bogus_id(args.url)))
+ results.append(("B", stage_b_abort_all_idle(args.url)))
+ results.append(("C", stage_c_abort_running(args.url, args.server_log_path)))
+ results.append(("D", stage_d_abort_by_id(args.url, args.server_log_path)))
+ results.append(("E", stage_e_health_after(args.url)))
+
+ print("\n" + "=" * 50)
+ all_ok = True
+ for name, passed in results:
+ tag = f"{GREEN}PASS{RESET}" if passed else f"{RED}FAIL{RESET}"
+ print(f" Stage {name}: {tag}")
+ all_ok = all_ok and passed
+ if not all_ok:
+ sys.exit(1)
+ print(f"\n{GREEN}ALL ABORT STAGES PASSED{RESET}")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/test/test_api/test_r3.py b/test/test_api/test_r3.py
new file mode 100644
index 0000000000..19dc910e61
--- /dev/null
+++ b/test/test_api/test_r3.py
@@ -0,0 +1,138 @@
+import sys
+import argparse
+import requests
+import base64
+import numpy as np
+
+
+def _check_prompt_logprobs(res, topk: int) -> bool:
+ prompt_token_ids = res.get("prompt_token_ids")
+ prompt_logprobs = res.get("prompt_logprobs")
+ if not isinstance(prompt_token_ids, list) or not isinstance(prompt_logprobs, list):
+ return False
+ if len(prompt_token_ids) != len(prompt_logprobs) or not prompt_logprobs or prompt_logprobs[0] is not None:
+ return False
+
+ expected_items = 1 if topk == 0 else topk
+ for position, position_logprobs in enumerate(prompt_logprobs[1:], start=1):
+ if len(position_logprobs) != expected_items:
+ return False
+ if topk == 0 and list(position_logprobs.keys()) != [str(prompt_token_ids[position])]:
+ return False
+ for item in position_logprobs.values():
+ if not np.isfinite(item["logprob"]):
+ return False
+ if topk == 0 and (not isinstance(item.get("rank"), int) or item["rank"] <= 0):
+ return False
+ if topk > 0 and any(item.get("rank") != index + 1 for index, item in enumerate(position_logprobs.values())):
+ return False
+ return True
+
+
+def test_routing_export(
+ url: str = "http://127.0.0.1:8000",
+ prompt_logprobs: int = 0,
+ timeout: int = 180,
+ max_new_tokens: int = 1,
+):
+ print(f"Testing routing export at {url}")
+ print(f"Requested prompt_logprobs: {prompt_logprobs}")
+ print("-" * 50)
+
+ try:
+ response = requests.post(
+ f"{url}/generate",
+ json={
+ "inputs": "你好,早上好!啊啊啊" * 5,
+ "parameters": {
+ "max_new_tokens": max_new_tokens,
+ "prompt_logprobs": prompt_logprobs,
+ "return_routed_experts": True,
+ # "repetition_penalty": 1.0,
+ },
+ },
+ timeout=timeout,
+ )
+ except requests.exceptions.ConnectionError:
+ print(f"ERROR: Cannot connect to server at {url}")
+ print(
+ "Make sure the LightLLM server is running with " "--enable_return_routed_experts --enable_prompt_logprobs"
+ )
+ return False
+ except requests.exceptions.Timeout:
+ print("ERROR: Request timed out")
+ return False
+
+ print(f"Status: {response.status_code}")
+
+ if response.status_code != 200:
+ print(f"ERROR: Request failed with status {response.status_code}")
+ print(f"Response: {response.text}")
+ return False
+
+ res = response.json()
+ print(f"Prompt tokens: {(res['prompt_token_ids'])}")
+ print(f"Prompt logprobs entries: {(res['prompt_logprobs'])}")
+ print(res["count_output_tokens"])
+ prompt_logprobs_ok = _check_prompt_logprobs(res, prompt_logprobs)
+
+ if "routed_experts" not in res or not res["routed_experts"]:
+ print("\nWARNING: No routed_experts in response.")
+ print("This could mean:")
+ print(" - The model is not a MoE model")
+ print(" - The server was not started with --enable_return_routed_experts")
+ print(" - The routing capture manager was not initialized")
+ return False
+ routing_info = res["routed_experts"]
+ shape = routing_info["shape"]
+ dtype_str = routing_info["dtype"]
+ dtype = np.dtype(dtype_str)
+ data = base64.b64decode(routing_info["data"])
+ routing_array = np.frombuffer(data, dtype=dtype).reshape(shape)
+
+ print(f"\n{'=' * 50}")
+ print("ROUTING CAPTURE SUCCESS!")
+ print(f"{'=' * 50}")
+ print(f"Shape: {shape}")
+ print(f"Dtype: {dtype}")
+ print(f"Num tokens: {shape[0]}")
+ print(f"Num MoE layers: {shape[1]}")
+ print(f"Top-K: {shape[2]}")
+
+ # Compute payload size savings
+ int32_size = np.prod(shape) * 4
+ actual_size = len(data)
+ savings = (1 - actual_size / int32_size) * 100
+ print(f"Payload: {actual_size} bytes (vs {int32_size} bytes with int32, {savings:.0f}% smaller)")
+
+ print(f"\nSample routing (first layer, first 5 tokens):")
+ num_tokens_to_show = min(shape[0], 5)
+ for i in range(num_tokens_to_show):
+ print(f" Token {i}: experts {routing_array[i, 0, :].tolist()}")
+
+ if np.all(routing_array == 0):
+ print("\nWARNING: All routing data is zeros. Capture may not be working correctly.")
+ return False
+
+ if not prompt_logprobs_ok:
+ return False
+
+ print("\nTest PASSED!")
+ return True
+
+
+if __name__ == "__main__":
+ parser = argparse.ArgumentParser(description="Test R3 routing export feature")
+ parser.add_argument("--url", default="http://127.0.0.1:8000", help="Server URL")
+ parser.add_argument("--prompt-logprobs", type=int, default=0, help="prompt_logprobs value to request")
+ parser.add_argument("--timeout", type=int, default=180, help="request timeout in seconds")
+ parser.add_argument("--max-new-tokens", type=int, default=1, help="max_new_tokens value to request")
+ args = parser.parse_args()
+
+ success = test_routing_export(
+ args.url,
+ prompt_logprobs=args.prompt_logprobs,
+ timeout=args.timeout,
+ max_new_tokens=args.max_new_tokens,
+ )
+ sys.exit(0 if success else 1)
diff --git a/test/test_api/test_rl_endpoints.py b/test/test_api/test_rl_endpoints.py
new file mode 100644
index 0000000000..91829530ec
--- /dev/null
+++ b/test/test_api/test_rl_endpoints.py
@@ -0,0 +1,349 @@
+"""
+Test release_memory_occupation / resume_memory_occupation / update_weights_from_tensor
+against a running lightllm server.
+
+Sequence:
+ 1. baseline generate (sanity)
+ 2. release_memory_occupation -> GPU memory should drop sharply
+ 3. resume_memory_occupation -> GPU memory should grow back
+ (without --enable_weight_cpu_backup the weight
+ memory is allocated empty, so generation right
+ after resume is expected to be garbage)
+ 4. update_weights_from_tensor (per-batch CUDA-IPC handoff) for every parameter
+ found on disk -> repopulate weights
+ 5. final generate -> should produce a sensible answer again
+
+The "trainer" runs in this same process: it holds tensors on a free GPU, serialises
+them via lightllm.utils.rl.serialization.MultiprocessingSerializer (CUDA IPC handles, not
+data), then asks the server to clone them into its weight buffers. No NCCL group
+is required, so this is safe to interrupt without leaving the server hung.
+
+Usage:
+ python test/test_api/test_rl_endpoints.py \
+ --url http://127.0.0.1:8000 \
+ --model_dir /nvme/models/Qwen3.5-35B-A3B \
+ --tp 4 \
+ --server_devices 0,1,2,3 \
+ --client_device 4
+
+Notes:
+ - This script must run on the same machine as the server (CUDA IPC).
+ - --server_devices are nvidia-smi GPU indices for the TP workers. If omitted,
+ the script infers them from the top --tp memory consumers before release.
+ - --client_device picks a free CUDA device for the in-process trainer; it is
+ independent from --server_devices and should not overlap the TP workers.
+"""
+
+import argparse
+import json
+import os
+import subprocess
+import sys
+import time
+from glob import glob
+from typing import Dict, List, Tuple
+
+# Make the repo importable when this script is invoked by path rather than -m.
+_REPO_ROOT = os.path.abspath(os.path.join(os.path.dirname(__file__), "..", ".."))
+if _REPO_ROOT not in sys.path:
+ sys.path.insert(0, _REPO_ROOT)
+
+import requests
+import torch
+from safetensors import safe_open
+
+from lightllm.utils.rl.serialization import MultiprocessingSerializer
+from lightllm.utils.rl.torch_cuda_ipc import monkey_patch_torch_reductions
+
+
+GREEN = "\033[32m"
+RED = "\033[31m"
+YELLOW = "\033[33m"
+RESET = "\033[0m"
+
+
+def banner(msg: str):
+ print(f"\n{YELLOW}=== {msg} ==={RESET}", flush=True)
+
+
+def ok(msg: str):
+ print(f" {GREEN}OK{RESET} {msg}", flush=True)
+
+
+def fail(msg: str):
+ print(f" {RED}FAIL{RESET} {msg}", flush=True)
+
+
+def gpu_mem_used_mib() -> List[int]:
+ out = subprocess.check_output(["nvidia-smi", "--query-gpu=memory.used", "--format=csv,noheader,nounits"]).decode()
+ return [int(x.strip()) for x in out.strip().splitlines()]
+
+
+def _select_gpu_mem(mem: List[int], devices: List[int]) -> List[int]:
+ return [mem[i] for i in devices]
+
+
+def _resolve_server_devices(server_devices: str, tp: int, mem: List[int]) -> List[int]:
+ if tp <= 0:
+ raise ValueError(f"--tp must be positive, got {tp}")
+ if tp > len(mem):
+ raise ValueError(f"--tp={tp} but nvidia-smi only returned {len(mem)} GPUs")
+
+ value = server_devices.strip()
+ if value.lower() == "auto":
+ return sorted(range(len(mem)), key=lambda i: mem[i], reverse=True)[:tp]
+
+ devices = [int(x.strip()) for x in value.split(",") if x.strip()]
+ if len(devices) != tp:
+ raise ValueError(f"--server_devices must contain exactly --tp entries; got {devices} for tp={tp}")
+ if len(set(devices)) != len(devices):
+ raise ValueError(f"--server_devices contains duplicates: {devices}")
+
+ bad = [i for i in devices if i < 0 or i >= len(mem)]
+ if bad:
+ raise ValueError(f"--server_devices contains invalid GPU indices {bad}; nvidia-smi returned {len(mem)} GPUs")
+ return devices
+
+
+def post(url: str, path: str, payload=None, timeout=600):
+ r = requests.post(url + path, json=payload or {}, timeout=timeout)
+ try:
+ body = r.json()
+ except Exception:
+ body = r.text
+ return r.status_code, body
+
+
+def generate(url: str, prompt: str, max_new_tokens: int = 16) -> str:
+ r = requests.post(
+ url + "/generate",
+ json={
+ "inputs": prompt,
+ "parameters": {"max_new_tokens": max_new_tokens, "do_sample": False},
+ },
+ timeout=120,
+ )
+ r.raise_for_status()
+ data = r.json()
+ if isinstance(data.get("generated_text"), list):
+ return data["generated_text"][0]
+ return data.get("generated_text", json.dumps(data))
+
+
+def looks_garbage(text: str) -> bool:
+ """Heuristic: post-resume text is usually a single repeated character (e.g. '!!!!')."""
+ s = text.strip()
+ if not s:
+ return True
+ return len(set(s)) == 1
+
+
+# ---------------- weight-update helpers (update_weights_from_tensor) ----------------
+
+
+def _list_safetensor_shards(model_dir: str) -> List[str]:
+ shards = sorted(glob(os.path.join(model_dir, "*.safetensors")))
+ if not shards:
+ raise RuntimeError(f"no .safetensors found under {model_dir}")
+ return shards
+
+
+def _send_update_from_tensor(
+ url: str,
+ serialized_per_rank: List[str],
+ flush_cache: bool = False,
+):
+ code, body = post(
+ url,
+ "/update_weights_from_tensor",
+ {
+ "serialized_named_tensors": serialized_per_rank,
+ "load_format": None,
+ "flush_cache": flush_cache,
+ "abort_all_requests": False,
+ },
+ timeout=600,
+ )
+ return code, body
+
+
+def update_weights_from_disk_via_tensor_api(
+ url: str,
+ model_dir: str,
+ tp: int,
+ client_device: int,
+ batch_per_request: int = 8,
+ flush_cache_at_end: bool = True,
+):
+ """
+ Acts as an in-process "trainer": loads every safetensor shard onto
+ cuda:client_device, then ships each batch of (name, tensor) to the server
+ via /update_weights_from_tensor. The server worker on each TP rank receives
+ a CUDA IPC handle, copies into its weight buffer.
+ """
+ banner("update_weights_from_tensor (CUDA IPC)")
+ # Server side patches its own copy; we patch ours so reductions can serialise
+ # CUDA tensors with UUID-based device addressing.
+ monkey_patch_torch_reductions()
+ torch.cuda.set_device(client_device)
+ device = f"cuda:{client_device}"
+
+ shards = _list_safetensor_shards(model_dir)
+ print(f" found {len(shards)} safetensor shards, batch_per_request={batch_per_request}", flush=True)
+
+ total_params = 0
+ total_bytes = 0
+ t0 = time.time()
+ for shard_idx, shard in enumerate(shards):
+ shard_t0 = time.time()
+ with safe_open(shard, framework="pt") as f:
+ keys = list(f.keys())
+ for i in range(0, len(keys), batch_per_request):
+ batch_keys = keys[i : i + batch_per_request]
+ # Load batch onto the client GPU. .contiguous() guarantees a
+ # whole-tensor allocation (safetensors slices are already
+ # contiguous, but this is cheap insurance).
+ tensors = [f.get_tensor(k).to(device).contiguous() for k in batch_keys]
+ named: List[Tuple[str, torch.Tensor]] = list(zip(batch_keys, tensors))
+
+ # Same payload to every TP rank — the server clones full
+ # tensors per rank and lets model.load_weights handle the TP
+ # sharding internally (matching how update_weights_from_*
+ # paths are written).
+ blob = MultiprocessingSerializer.serialize(named, output_str=True)
+ serialized_per_rank = [blob] * tp
+
+ # Last batch flushes the prefix cache so old KV from the
+ # previous weight version cannot poison subsequent gens.
+ is_last = (shard_idx == len(shards) - 1) and (i + batch_per_request >= len(keys))
+ code, body = _send_update_from_tensor(
+ url,
+ serialized_per_rank,
+ flush_cache=(flush_cache_at_end and is_last),
+ )
+ if code != 200:
+ fail(f"update batch failed: {code} {body}")
+ raise RuntimeError(f"update batch failed: {code} {body}")
+ total_params += len(batch_keys)
+ total_bytes += sum(t.numel() * t.element_size() for t in tensors)
+ # Free client-side memory before next batch — the worker has
+ # already cloned the data by the time post() returned.
+ for t in tensors:
+ del t
+ del tensors, named
+ torch.cuda.empty_cache()
+
+ print(
+ f" shard {shard_idx+1}/{len(shards)} done "
+ f"(+{len(keys)} tensors, {time.time()-shard_t0:.1f}s, "
+ f"running total {total_params} params, {total_bytes/1e9:.1f} GB)",
+ flush=True,
+ )
+
+ dt = time.time() - t0
+ ok(f"streamed {total_params} params, {total_bytes/1e9:.1f} GB in {dt:.1f}s")
+
+
+# ---------------- main flow ----------------
+
+
+def main():
+ ap = argparse.ArgumentParser()
+ ap.add_argument("--url", default="http://127.0.0.1:8000")
+ ap.add_argument("--model_dir", required=True)
+ ap.add_argument("--tp", type=int, required=True)
+ ap.add_argument(
+ "--server_devices",
+ default="auto",
+ help="comma-separated nvidia-smi GPU indices used by the server, or 'auto' to infer from memory usage",
+ )
+ ap.add_argument(
+ "--client_device",
+ type=int,
+ default=2,
+ help="GPU index for the in-process trainer; must differ from TP worker GPUs",
+ )
+ ap.add_argument("--prompt", default="The capital of France is")
+ ap.add_argument("--max_new_tokens", type=int, default=16)
+ ap.add_argument("--batch_per_request", type=int, default=8)
+ ap.add_argument("--skip_update", action="store_true", help="run only release/resume, skip the update_weights phase")
+ args = ap.parse_args()
+
+ # ---------------- stage 1: baseline ----------------
+ banner("baseline generate")
+ base_text = generate(args.url, args.prompt, args.max_new_tokens)
+ print(f" prompt : {args.prompt!r}")
+ print(f" generated: {base_text!r}")
+ ok("baseline generated")
+
+ # ---------------- stage 2: release ----------------
+ banner("release_memory_occupation")
+ before = gpu_mem_used_mib()
+ try:
+ server_devices = _resolve_server_devices(args.server_devices, args.tp, before)
+ except ValueError as e:
+ fail(str(e))
+ sys.exit(1)
+ print(f" server GPUs : {server_devices}")
+ print(f" GPU mem before: {_select_gpu_mem(before, server_devices)}")
+ code, body = post(args.url, "/release_memory_occupation", {})
+ print(f" resp: {code} {body}")
+ if code != 200:
+ fail("release failed")
+ sys.exit(1)
+ time.sleep(2)
+ after = gpu_mem_used_mib()
+ print(f" GPU mem after : {_select_gpu_mem(after, server_devices)}")
+ drop = sum(_select_gpu_mem(before, server_devices)) - sum(_select_gpu_mem(after, server_devices))
+ if drop < 10_000:
+ fail(f"release did not free much memory (delta={drop} MiB)")
+ sys.exit(1)
+ ok(f"release freed ~{drop} MiB on TP GPUs")
+
+ # ---------------- stage 3: resume ----------------
+ banner("resume_memory_occupation")
+ code, body = post(args.url, "/resume_memory_occupation", {})
+ print(f" resp: {code} {body}")
+ if code != 200:
+ fail("resume failed")
+ sys.exit(1)
+ time.sleep(2)
+ print(f" GPU mem after : {_select_gpu_mem(gpu_mem_used_mib(), server_devices)}")
+ ok("resume returned success")
+
+ banner("post-resume generate (likely garbage without weight cpu backup)")
+ text_after_resume = generate(args.url, args.prompt, args.max_new_tokens)
+ print(f" generated: {text_after_resume!r} garbage_heuristic={looks_garbage(text_after_resume)}")
+
+ if args.skip_update:
+ ok("done (skipped update_weights stage)")
+ return
+
+ # ---------------- stage 4: update_weights_from_tensor ----------------
+ update_weights_from_disk_via_tensor_api(
+ url=args.url,
+ model_dir=args.model_dir,
+ tp=args.tp,
+ client_device=args.client_device,
+ batch_per_request=args.batch_per_request,
+ flush_cache_at_end=True,
+ )
+
+ # ---------------- stage 5: final generate ----------------
+ banner("final generate (after weight reload)")
+ final_text = generate(args.url, args.prompt, args.max_new_tokens)
+ print(f" prompt : {args.prompt!r}")
+ print(f" generated: {final_text!r}")
+ if looks_garbage(final_text):
+ fail("final generation still looks like garbage; weight update did not stick")
+ sys.exit(1)
+ if final_text.strip() == base_text.strip():
+ ok("final output matches baseline exactly")
+ else:
+ ok("final output is sensible (differs from baseline but not garbage)")
+
+ print(f"\n{GREEN}ALL STAGES PASSED{RESET}")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/unit_tests/__init__.py b/unit_tests/__init__.py
new file mode 100644
index 0000000000..e69de29bb2
diff --git a/unit_tests/common/__init__.py b/unit_tests/common/__init__.py
new file mode 100644
index 0000000000..e69de29bb2
diff --git a/unit_tests/common/basemodel/__init__.py b/unit_tests/common/basemodel/__init__.py
new file mode 100644
index 0000000000..e69de29bb2
diff --git a/unit_tests/common/basemodel/test_moe_route_info_manager.py b/unit_tests/common/basemodel/test_moe_route_info_manager.py
new file mode 100644
index 0000000000..de2bc7a857
--- /dev/null
+++ b/unit_tests/common/basemodel/test_moe_route_info_manager.py
@@ -0,0 +1,263 @@
+from types import SimpleNamespace
+
+import numpy as np
+import pytest
+import torch
+
+from lightllm.common.basemodel.infer_struct import InferStateInfo
+from lightllm.common.basemodel.moe_route_info_manager import MoeRouteInfoManager, get_moe_capture_callback
+
+
+def _skip_without_cuda():
+ if not torch.cuda.is_available():
+ pytest.skip("CUDA is required for moe route info capture.")
+
+
+def test_get_moe_capture_callback_uses_global_manager(monkeypatch):
+ calls = []
+
+ class _Manager:
+ def is_buffer_initialized(self):
+ return True
+
+ def get_moe_capture_callback(self, layer_index, mem_indexes):
+ calls.append((layer_index, mem_indexes))
+ return "callback"
+
+ mem_indexes = object()
+ infer_state = SimpleNamespace(mem_index=mem_indexes)
+ monkeypatch.setattr(MoeRouteInfoManager, "_instance", _Manager())
+
+ assert get_moe_capture_callback(infer_state, 7) == "callback"
+ assert calls == [(7, mem_indexes)]
+
+
+class TestMoeRouteInfoManager:
+ def test_capture_and_extract_basic(self):
+ _skip_without_cuda()
+
+ manager = MoeRouteInfoManager(
+ num_moe_layers=4,
+ topk=8,
+ dtype_id=1,
+ )
+ manager.init_capture_buffer(kv_cache_size=1024)
+ mem_indexes = torch.arange(100, 110, device="cuda")
+ expected = np.zeros((10, 4, 8), dtype=np.uint8)
+
+ for layer_idx in range(4):
+ topk_ids = torch.randint(0, 64, (10, 8), device="cuda")
+ manager.get_moe_capture_callback(layer_idx, mem_indexes)(topk_ids)
+ expected[:, layer_idx, :] = topk_ids.cpu().numpy().astype(np.uint8)
+
+ result = manager.extract(mem_indexes)
+ assert result.shape == (10, 4, 8)
+ assert result.dtype == np.uint8
+ np.testing.assert_array_equal(result, expected)
+
+ def test_capture_writes_to_correct_kv_positions(self):
+ _skip_without_cuda()
+
+ manager = MoeRouteInfoManager(
+ num_moe_layers=2,
+ topk=4,
+ dtype_id=1,
+ )
+ manager.init_capture_buffer(kv_cache_size=256)
+ mem_indexes = torch.tensor([10, 50, 200], device="cuda")
+ topk_ids = torch.tensor([[1, 2, 3, 4], [5, 6, 7, 8], [9, 10, 11, 12]], device="cuda")
+ topk_ids_layer1 = topk_ids + 20
+
+ manager.get_moe_capture_callback(0, mem_indexes)(topk_ids)
+ manager.get_moe_capture_callback(1, mem_indexes)(topk_ids_layer1)
+
+ result = manager.extract(mem_indexes)
+ np.testing.assert_array_equal(result[:, 0, :], topk_ids.cpu().numpy().astype(np.uint8))
+ np.testing.assert_array_equal(result[:, 1, :], topk_ids_layer1.cpu().numpy().astype(np.uint8))
+
+ def test_capture_maps_transformer_layer_num_to_routing_slot(self):
+ _skip_without_cuda()
+
+ manager = MoeRouteInfoManager(
+ num_moe_layers=2,
+ topk=2,
+ dtype_id=1,
+ layer_index_to_moe_index={3: 0, 7: 1},
+ )
+ manager.init_capture_buffer(kv_cache_size=256)
+ mem_indexes = torch.tensor([10, 11], device="cuda")
+ ids_layer3 = torch.tensor([[1, 2], [3, 4]], device="cuda")
+ ids_layer7 = torch.tensor([[5, 6], [7, 8]], device="cuda")
+
+ manager.get_moe_capture_callback(3, mem_indexes)(ids_layer3)
+ manager.get_moe_capture_callback(7, mem_indexes)(ids_layer7)
+ assert manager.get_moe_capture_callback(4, mem_indexes) is None
+
+ result = manager.extract(mem_indexes)
+ np.testing.assert_array_equal(result[:, 0, :], ids_layer3.cpu().numpy().astype(np.uint8))
+ np.testing.assert_array_equal(result[:, 1, :], ids_layer7.cpu().numpy().astype(np.uint8))
+
+ def test_capture_rejects_unexpected_topk_width(self):
+ _skip_without_cuda()
+
+ manager = MoeRouteInfoManager(
+ num_moe_layers=1,
+ topk=2,
+ dtype_id=1,
+ )
+ manager.init_capture_buffer(kv_cache_size=256)
+ mem_indexes = torch.tensor([10, 11], device="cuda")
+ topk_ids_with_shared_expert = torch.tensor([[1, 2, 32], [3, 4, 32]], device="cuda")
+
+ with pytest.raises(AssertionError):
+ manager.get_moe_capture_callback(0, mem_indexes)(topk_ids_with_shared_expert)
+
+ def test_cuda_graph_replay_uses_copied_mem_indexes(self):
+ _skip_without_cuda()
+
+ class _NoopDecodeState:
+ def copy_for_decode_cuda_graph(self, other):
+ return
+
+ manager = MoeRouteInfoManager(
+ num_moe_layers=1,
+ topk=2,
+ dtype_id=1,
+ )
+ manager.init_capture_buffer(kv_cache_size=256)
+ graph_infer_state = InferStateInfo()
+ graph_infer_state.decode_att_state = _NoopDecodeState()
+ graph_infer_state.mem_index = torch.tensor([10, 11], device="cuda")
+ moe_capture_callback = manager.get_moe_capture_callback(0, graph_infer_state.mem_index)
+ topk_ids = torch.tensor([[1, 2], [3, 4]], device="cuda")
+
+ graph = torch.cuda.CUDAGraph()
+ torch.cuda.synchronize()
+ with torch.cuda.graph(graph):
+ moe_capture_callback(topk_ids)
+
+ graph.replay()
+ result = manager.extract(graph_infer_state.mem_index)
+ np.testing.assert_array_equal(result[:, 0, :], topk_ids.cpu().numpy().astype(np.uint8))
+
+ new_infer_state = InferStateInfo()
+ new_infer_state.decode_att_state = _NoopDecodeState()
+ new_infer_state.mem_index = torch.tensor([20, 21], device="cuda")
+ new_topk_ids = torch.tensor([[5, 6], [7, 8]], device="cuda")
+ graph_infer_state.copy_for_cuda_graph(new_infer_state)
+ topk_ids.copy_(new_topk_ids)
+
+ graph.replay()
+
+ result = manager.extract(new_infer_state.mem_index)
+ np.testing.assert_array_equal(result[:, 0, :], new_topk_ids.cpu().numpy().astype(np.uint8))
+
+ def test_microbatch_isolation(self):
+ _skip_without_cuda()
+
+ manager = MoeRouteInfoManager(
+ num_moe_layers=1,
+ topk=4,
+ dtype_id=1,
+ )
+ manager.init_capture_buffer(kv_cache_size=256)
+ mem0 = torch.tensor([10, 11], device="cuda")
+ mem1 = torch.tensor([20, 21], device="cuda")
+ ids_0 = torch.ones((2, 4), dtype=torch.int64, device="cuda")
+ ids_1 = torch.ones((2, 4), dtype=torch.int64, device="cuda") * 2
+
+ capture0 = manager.get_moe_capture_callback(0, mem0)
+ capture1 = manager.get_moe_capture_callback(0, mem1)
+ capture0(ids_0)
+ capture1(ids_1)
+
+ result0 = manager.extract(mem0)
+ result1 = manager.extract(mem1)
+ assert result0[0, 0, 0] == 1
+ assert result1[0, 0, 0] == 2
+
+ def test_dtype_selection_uint8(self):
+ manager = MoeRouteInfoManager(
+ num_moe_layers=1,
+ topk=2,
+ dtype_id=1,
+ )
+ assert manager.get_torch_dtype() == torch.uint8
+ assert manager.get_np_dtype() == np.uint8
+ assert manager.dtype_id == 1
+ assert not manager.is_buffer_initialized()
+
+ def test_dtype_selection_int16(self):
+ manager = MoeRouteInfoManager(
+ num_moe_layers=1,
+ topk=2,
+ dtype_id=2,
+ )
+ assert manager.get_torch_dtype() == torch.int16
+ assert manager.get_np_dtype() == np.int16
+ assert manager.dtype_id == 2
+ assert not manager.is_buffer_initialized()
+
+ def test_extract_preserves_uint8_values(self):
+ _skip_without_cuda()
+
+ manager = MoeRouteInfoManager(
+ num_moe_layers=1,
+ topk=4,
+ dtype_id=1,
+ )
+ manager.init_capture_buffer(kv_cache_size=64)
+ mem_indexes = torch.tensor([0, 1, 2], device="cuda")
+ topk_ids = torch.tensor([[10, 20, 30, 40], [50, 60, 63, 1], [0, 5, 255, 3]], device="cuda")
+
+ manager.get_moe_capture_callback(0, mem_indexes)(topk_ids)
+
+ result = manager.extract(mem_indexes)
+ np.testing.assert_array_equal(result[:, 0, :], topk_ids.cpu().numpy().astype(np.uint8))
+
+ def test_routing_buffer_and_pointer_shape(self):
+ _skip_without_cuda()
+
+ manager = MoeRouteInfoManager(
+ num_moe_layers=48,
+ topk=8,
+ dtype_id=1,
+ )
+ assert not manager.is_buffer_initialized()
+ manager.init_capture_buffer(kv_cache_size=2048)
+ assert manager.routing_buffer.shape == (2048, 48, 8)
+ assert manager.routing_buffer.dtype == torch.uint8
+ assert manager.routing_buffer.device.type == "cpu"
+ assert manager.routing_buffer.is_pinned()
+ assert manager.routing_buffer_ptr.shape == (1,)
+ assert manager.routing_buffer_ptr.dtype == torch.uint64
+ assert manager.routing_buffer_ptr.device.type == "cuda"
+
+ def test_partial_token_capture(self):
+ _skip_without_cuda()
+
+ manager = MoeRouteInfoManager(
+ num_moe_layers=1,
+ topk=2,
+ dtype_id=1,
+ )
+ manager.init_capture_buffer(kv_cache_size=128)
+ mem_indexes = torch.tensor([10, 11, 12, 13, 14], device="cuda")
+ topk_ids = torch.tensor([[1, 2], [3, 4], [5, 6]], device="cuda")
+
+ manager.get_moe_capture_callback(0, mem_indexes[:3])(topk_ids)
+
+ result_written = manager.extract(mem_indexes[:3])
+ np.testing.assert_array_equal(result_written[:, 0, :], topk_ids.cpu().numpy().astype(np.uint8))
+
+ result_unwritten = manager.extract(mem_indexes[3:])
+ np.testing.assert_array_equal(result_unwritten[:, 0, :], np.zeros((2, 2), dtype=np.uint8))
+
+ def test_phase1_does_not_allocate_capture_buffer(self):
+ manager = MoeRouteInfoManager(
+ num_moe_layers=4,
+ topk=8,
+ dtype_id=1,
+ )
+ assert not manager.is_buffer_initialized()
+ assert manager.routing_buffer_ptr is None
diff --git a/unit_tests/common/basemodel/triton_kernel/test_routing_capture.py b/unit_tests/common/basemodel/triton_kernel/test_routing_capture.py
new file mode 100644
index 0000000000..2bd9bdf06c
--- /dev/null
+++ b/unit_tests/common/basemodel/triton_kernel/test_routing_capture.py
@@ -0,0 +1,133 @@
+import pytest
+import torch
+
+from lightllm.common.basemodel.triton_kernel.routing_capture import scatter_routing_topk_to_cpu
+
+
+def _skip_without_cuda():
+ if not torch.cuda.is_available():
+ pytest.skip("CUDA is required for Triton kernels.")
+
+
+@pytest.mark.parametrize("dtype,dtype_id,max_value", [(torch.uint8, 1, 255), (torch.int16, 2, 1024)])
+def test_scatter_routing_topk_to_pinned_cpu(dtype, dtype_id, max_value):
+ _skip_without_cuda()
+
+ num_tokens = 5
+ num_moe_layers = 3
+ topk = 4
+ kv_cache_size = 32
+ moe_layer_index = 1
+
+ topk_ids = torch.randint(0, max_value, (num_tokens, topk), dtype=torch.int64, device="cuda")
+ mem_indexes = torch.tensor([17, 3, 29, 8, 11], dtype=torch.int32, device="cuda")
+ routing_buffer = torch.zeros(
+ (kv_cache_size, num_moe_layers, topk),
+ dtype=dtype,
+ device="cpu",
+ pin_memory=True,
+ )
+ routing_buffer_ptr = torch.tensor([routing_buffer.data_ptr()], dtype=torch.uint64, device="cuda")
+
+ scatter_routing_topk_to_cpu(
+ topk_ids=topk_ids,
+ mem_indexes=mem_indexes,
+ routing_buffer_ptr=routing_buffer_ptr,
+ moe_layer_index=moe_layer_index,
+ num_moe_layers=num_moe_layers,
+ topk=topk,
+ dtype_id=dtype_id,
+ )
+ torch.cuda.synchronize()
+
+ expected = torch.zeros_like(routing_buffer)
+ expected[mem_indexes.cpu().long(), moe_layer_index, :] = topk_ids.cpu().to(dtype)
+ assert torch.equal(routing_buffer, expected)
+
+
+def test_scatter_routing_topk_respects_layer_index():
+ _skip_without_cuda()
+
+ num_tokens = 3
+ num_moe_layers = 2
+ topk = 2
+ kv_cache_size = 16
+
+ topk_ids = torch.arange(num_tokens * topk, dtype=torch.int64, device="cuda").view(num_tokens, topk)
+ mem_indexes = torch.tensor([10, 4, 13], dtype=torch.int64, device="cuda")
+ routing_buffer = torch.zeros(
+ (kv_cache_size, num_moe_layers, topk),
+ dtype=torch.int16,
+ device="cpu",
+ pin_memory=True,
+ )
+ routing_buffer_ptr = torch.tensor([routing_buffer.data_ptr()], dtype=torch.uint64, device="cuda")
+
+ scatter_routing_topk_to_cpu(
+ topk_ids=topk_ids,
+ mem_indexes=mem_indexes,
+ routing_buffer_ptr=routing_buffer_ptr,
+ moe_layer_index=1,
+ num_moe_layers=num_moe_layers,
+ topk=topk,
+ dtype_id=2,
+ )
+ torch.cuda.synchronize()
+
+ expected = torch.zeros_like(routing_buffer)
+ expected[mem_indexes.cpu(), 1, :] = topk_ids.cpu().to(torch.int16)
+ assert torch.equal(routing_buffer, expected)
+
+
+def test_scatter_routing_topk_is_cuda_graph_capturable():
+ _skip_without_cuda()
+
+ num_tokens = 4
+ num_moe_layers = 2
+ topk = 3
+ kv_cache_size = 16
+
+ topk_ids = torch.arange(num_tokens * topk, dtype=torch.int64, device="cuda").view(num_tokens, topk)
+ mem_indexes = torch.tensor([2, 4, 6, 8], dtype=torch.int32, device="cuda")
+ routing_buffer = torch.zeros(
+ (kv_cache_size, num_moe_layers, topk),
+ dtype=torch.uint8,
+ device="cpu",
+ pin_memory=True,
+ )
+ routing_buffer_ptr = torch.tensor([routing_buffer.data_ptr()], dtype=torch.uint64, device="cuda")
+
+ scatter_routing_topk_to_cpu(
+ topk_ids=topk_ids,
+ mem_indexes=mem_indexes,
+ routing_buffer_ptr=routing_buffer_ptr,
+ moe_layer_index=0,
+ num_moe_layers=num_moe_layers,
+ topk=topk,
+ dtype_id=1,
+ )
+ torch.cuda.synchronize()
+ routing_buffer.zero_()
+
+ graph = torch.cuda.CUDAGraph()
+ with torch.cuda.graph(graph):
+ scatter_routing_topk_to_cpu(
+ topk_ids=topk_ids,
+ mem_indexes=mem_indexes,
+ routing_buffer_ptr=routing_buffer_ptr,
+ moe_layer_index=0,
+ num_moe_layers=num_moe_layers,
+ topk=topk,
+ dtype_id=1,
+ )
+
+ graph.replay()
+ torch.cuda.synchronize()
+
+ expected = torch.zeros_like(routing_buffer)
+ expected[mem_indexes.cpu(), 0, :] = topk_ids.cpu().to(torch.uint8)
+ assert torch.equal(routing_buffer, expected)
+
+
+if __name__ == "__main__":
+ pytest.main([__file__])
diff --git a/unit_tests/server/core/objs/test_req.py b/unit_tests/server/core/objs/test_req.py
index 1c946531c1..584f421928 100644
--- a/unit_tests/server/core/objs/test_req.py
+++ b/unit_tests/server/core/objs/test_req.py
@@ -1,6 +1,7 @@
import pytest
import easydict
from lightllm.server.core.objs.req import Req, TokenHealingReq, ChunkedPrefillReq, SamplingParams
+from lightllm.server.core.objs.token_metadata import ReqFinalTokenMetadata
from lightllm.utils.envs_utils import set_env_start_args
@@ -14,6 +15,7 @@ def setup_module_env():
"llm_decode_att_backend": ["None"],
"cpu_cache_token_page_size": 256,
"enable_cpu_cache": False,
+ "model_dir": "",
}
)
)
@@ -40,6 +42,23 @@ def test_get_used_tokens(req):
assert req.get_used_tokens() == 5
+def test_final_token_metadata_read_returns_actual_prompt_tokens(req):
+ req.sample_params.prompt_logprobs = 0
+ req.shm_logprobs.arr["logprob"][1] = -0.5
+ req.shm_logprobs.arr["logprob"][2] = -1.25
+ req.shm_logprobs.arr["rank"][1] = 315
+ req.shm_logprobs.arr["rank"][2] = 4
+
+ metadata = ReqFinalTokenMetadata(req).read()
+
+ assert metadata["prompt_token_ids"] == [1, 2, 3]
+ assert metadata["prompt_logprobs"] == [
+ None,
+ {2: {"logprob": -0.5, "rank": 315, "decoded_token": None}},
+ {3: {"logprob": -1.25, "rank": 4, "decoded_token": None}},
+ ]
+
+
def test_token_healing_req_post_init():
token_healing_req = TokenHealingReq()
token_healing_req.init(1, [1, 2, 3, 4], {"max_new_tokens": 1}, None)
@@ -56,6 +75,12 @@ def test_token_healing_req_post_init():
def test_finish_status(req):
req.finish_status.set_status(req.finish_status.FINISHED_STOP)
assert req.finish_status.is_finished()
+ assert req.finish_status.get_finish_reason() == "stop"
+
+ req.finish_status.set_status(req.finish_status.FINISHED_ERROR)
+ assert req.finish_status.is_finished()
+ assert req.finish_status.is_finished_error()
+ assert req.finish_status.get_finish_reason() == "error"
if __name__ == "__main__":
diff --git a/unit_tests/server/router/dynamic_prompt/test_radix_cache.py b/unit_tests/server/router/dynamic_prompt/test_radix_cache.py
index 605433e9d8..dfeda0b6f7 100644
--- a/unit_tests/server/router/dynamic_prompt/test_radix_cache.py
+++ b/unit_tests/server/router/dynamic_prompt/test_radix_cache.py
@@ -230,5 +230,32 @@ def test_case9():
assert torch.equal(unmerged_node_d.token_id_key, torch.tensor([6], dtype=torch.int64))
+def test_case10():
+ """
+ 测试场景:测试 flush_cache 函数
+ """
+ print("\nTest Case 10: Testing flush_cache function\n")
+ tree = RadixCache("unique_name", 100, 0)
+ tree.insert(torch.tensor([1, 2, 3], dtype=torch.int64))
+ tree.insert(torch.tensor([1, 2, 3, 4, 5], dtype=torch.int64))
+ tree_node, size, values = tree.match_prefix(
+ torch.tensor([1, 2, 3], dtype=torch.int64, device="cpu"), update_refs=True
+ )
+ assert tree_node is not None
+ assert size == 3
+ tree.flush_cache()
+ tree_node, size, values = tree.match_prefix(
+ torch.tensor([1, 2, 3], dtype=torch.int64, device="cpu"), update_refs=True
+ )
+ assert tree_node is None
+ assert size == 0
+ assert tree.get_tree_total_tokens_num() == 0
+ assert tree.get_refed_tokens_num() == 0
+ assert len(tree.root_node.children) == 0
+ assert tree.root_node.token_id_key.numel() == 0
+ assert tree.root_node.token_mem_index_value.numel() == 0
+ assert tree.root_node.ref_counter == 1
+
+
if __name__ == "__main__":
pytest.main()