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()