diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index e80f2b552f..bebf2671cc 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -23,6 +23,9 @@ from lightllm.common.basemodel.prefill_cuda_graph import PrefillCudaGraph from lightllm.common.quantization import Quantcfg from lightllm.common.basemodel.triton_kernel.gather_token_id import gather_token, gather_token_prefill_decode_mixed +from lightllm.common.basemodel.triton_kernel.post_process.vocab_parallel_topk import ( + is_vocab_parallel_topk_enabled, +) from lightllm.utils.log_utils import init_logger from lightllm.utils.dist_utils import get_dp_world_size from lightllm.utils.profile_max_tokens import profile_mtp_weight_memory @@ -378,12 +381,22 @@ def forward(self, model_input: ModelInput): else: return self._decode(model_input) + def _is_cuda_graph_output_compatible(self, *model_inputs: ModelInput) -> bool: + """Whether inputs match the dense/sparse contract captured at startup.""" + + return ( + self.is_mtp_draft_model + or not is_vocab_parallel_topk_enabled() + or all(model_input.use_vocab_parallel_topk for model_input in model_inputs) + ) + def _create_inferstate(self, model_input: ModelInput, microbatch_index: int = 0): infer_state = self.infer_state_class() infer_state.hidden_collector = self.hidden_collector_prototype.new_instance() infer_state.input_ids = model_input.input_ids infer_state.is_prefill = model_input.is_prefill infer_state.return_all_prompt_logics = self.return_all_prompt_logics + infer_state.use_vocab_parallel_topk = self.is_mtp_draft_model or model_input.use_vocab_parallel_topk infer_state.batch_size = model_input.batch_size infer_state.total_token_num = model_input.total_token_num infer_state.max_q_seq_len = model_input.max_q_seq_len @@ -534,6 +547,8 @@ def _create_unpad_decode_model_output(self, model_output: ModelOutput, origin_ba return model_output new_model_output = copy.copy(model_output) new_model_output.logits = new_model_output.logits[0:origin_batch_size] + if new_model_output.logits_token_ids is not None: + new_model_output.logits_token_ids = new_model_output.logits_token_ids[0:origin_batch_size] new_model_output.mtp_collector = model_output.mtp_collector.unpad_decode( padded_batch_size=padded_batch_size, origin_batch_size=origin_batch_size, @@ -546,6 +561,8 @@ def _create_unpad_prefill_model_output( new_model_output = copy.copy(padded_model_output) # logits 始终只对应每个请求最后一个位置,移除 padding 的 req 对应的行。 new_model_output.logits = new_model_output.logits[0:origin_batch_size] + if new_model_output.logits_token_ids is not None: + new_model_output.logits_token_ids = new_model_output.logits_token_ids[0:origin_batch_size] new_model_output.mtp_collector = padded_model_output.mtp_collector.unpad_prefill( origin_handle_token_num=origin_handle_token_num ) @@ -643,9 +660,13 @@ def _decode( # CUDA Graph 可能继续向上对齐 batch size,并因此加入 seq_len=2 的 # dummy request。先用最终可能出现的 KV 长度判断 graph,再统一 padding 一次。 infer_max_kv_seq_len = max(2, model_input.max_kv_seq_len) - use_cuda_graph = self.graph is not None and self.graph.can_run( - batch_size=infer_batch_size, - max_len_in_batch=infer_max_kv_seq_len, + use_cuda_graph = ( + self._is_cuda_graph_output_compatible(model_input) + and self.graph is not None + and self.graph.can_run( + batch_size=infer_batch_size, + max_len_in_batch=infer_max_kv_seq_len, + ) ) need_capture = False if use_cuda_graph: @@ -678,7 +699,6 @@ def _decode( @final def _context_forward(self, infer_state: InferStateInfo): - input_embs = self.pre_infer.context_forward(infer_state.input_ids, infer_state, self.pre_post_weight) if self.args.enable_dp_prefill_balance: assert not self.args.enable_prefill_cudagraph, "not support now" @@ -737,6 +757,7 @@ def prefill_func(input_tensors, _infer_state): hidden_collector.add_final_hidden(last_input_embs) model_output = ModelOutput( logits=predict_logits.contiguous(), + logits_token_ids=infer_state.logits_token_ids, mtp_collector=infer_state.hidden_collector.finish_output(infer_state=infer_state), prompt_logics=infer_state.prompt_logics, ) @@ -766,6 +787,7 @@ def _token_forward(self, infer_state: InferStateInfo): hidden_collector.add_final_hidden(last_input_embs) model_output = ModelOutput( logits=predict_logits.contiguous(), + logits_token_ids=infer_state.logits_token_ids, mtp_collector=infer_state.hidden_collector.finish_output(infer_state=infer_state), ) @@ -897,7 +919,11 @@ def _microbatch_overlap_decode_cuda(self, model_input0: ModelInput, model_input1 infer_batch_size = max(1, origin_batch_size0, origin_batch_size1) infer_batch_size = triton.cdiv(infer_batch_size, self.tp_world_size_) * self.tp_world_size_ - if self.graph is not None and self.graph.can_run(infer_batch_size, max_len_in_batch): + if ( + self._is_cuda_graph_output_compatible(model_input0, model_input1) + and self.graph is not None + and self.graph.can_run(infer_batch_size, max_len_in_batch) + ): infer_batch_size = self.graph.find_closest_graph_batch_size(infer_batch_size) need_capture = self.graph.need_capture(infer_batch_size) padded_model_input0 = self._create_padded_decode_model_input(model_input0, infer_batch_size) @@ -1020,11 +1046,13 @@ def _overlap_tpsp_context_forward(self, infer_state: InferStateInfo, infer_state hidden_collector1.add_final_hidden(last_input_embs1) model_output = ModelOutput( logits=predict_logits.contiguous(), + logits_token_ids=infer_state.logits_token_ids, mtp_collector=infer_state.hidden_collector.finish_output(infer_state=infer_state), prompt_logics=infer_state.prompt_logics, ) model_output1 = ModelOutput( logits=predict_logits1.contiguous(), + logits_token_ids=infer_state1.logits_token_ids, mtp_collector=infer_state1.hidden_collector.finish_output(infer_state=infer_state1), prompt_logics=infer_state1.prompt_logics, ) @@ -1069,10 +1097,12 @@ def _overlap_tpsp_token_forward(self, infer_state: InferStateInfo, infer_state1: hidden_collector1.add_final_hidden(last_input_embs1) model_output = ModelOutput( logits=predict_logits.contiguous(), + logits_token_ids=infer_state.logits_token_ids, mtp_collector=infer_state.hidden_collector.finish_output(infer_state=infer_state), ) model_output1 = ModelOutput( logits=predict_logits1.contiguous(), + logits_token_ids=infer_state1.logits_token_ids, mtp_collector=infer_state1.hidden_collector.finish_output(infer_state=infer_state1), ) diff --git a/lightllm/common/basemodel/batch_objs.py b/lightllm/common/basemodel/batch_objs.py index ae645d4b7b..77ed1d400d 100644 --- a/lightllm/common/basemodel/batch_objs.py +++ b/lightllm/common/basemodel/batch_objs.py @@ -55,6 +55,10 @@ class ModelInput: # 的 draft 模型的输入 mtp_draft_input_hiddens: Optional[torch.Tensor] = None + # The router sets this only when a target-model batch can sample directly + # from sparse candidates. Draft models always enable the same output form. + use_vocab_parallel_topk: bool = False + def to_cuda(self): self.check_input() @@ -200,10 +204,49 @@ class ModelOutput: # 需要返回 prompt logprobs 信息时才会非空。 prompt_logics: Optional[torch.Tensor] = None + # Sparse vocab-parallel outputs map every candidate column back to its + # global token id. None means logits are dense and column indexes are ids. + logits_token_ids: Optional[torch.Tensor] = None + def __post_init__(self) -> None: if self.mtp_collector is None: self.mtp_collector = ModelMtpOutputCollector() + if self.logits_token_ids is not None: + assert self.logits.ndim == 2 + assert self.logits_token_ids.shape == self.logits.shape + assert self.logits_token_ids.dtype in (torch.int32, torch.int64) + assert self.logits_token_ids.device == self.logits.device def to_no_ref_tensor(self): self.logits = tensor_to_no_ref_tensor(self.logits) + if self.logits_token_ids is not None: + self.logits_token_ids = tensor_to_no_ref_tensor(self.logits_token_ids) self.mtp_collector.to_no_ref_tensor() + + @property + def has_vocab_parallel_logits(self) -> bool: + return self.logits_token_ids is not None + + def index_select_logits_rows(self, index: torch.Tensor) -> "ModelOutput": + """Select logit rows without dropping their vocabulary metadata.""" + + return ModelOutput( + logits=self.logits.index_select(0, index), + logits_token_ids=( + self.logits_token_ids.index_select(0, index) if self.logits_token_ids is not None else None + ), + ) + + @classmethod + def concat_logits_rows(cls, outputs: List["ModelOutput"]) -> "ModelOutput": + """Concatenate outputs that share the same dense or sparse layout.""" + + assert outputs + has_vocab_parallel_logits = outputs[0].has_vocab_parallel_logits + assert all(output.has_vocab_parallel_logits == has_vocab_parallel_logits for output in outputs) + return cls( + logits=torch.cat([output.logits for output in outputs], dim=0), + logits_token_ids=( + torch.cat([output.logits_token_ids for output in outputs], dim=0) if has_vocab_parallel_logits else None + ), + ) diff --git a/lightllm/common/basemodel/cuda_graph.py b/lightllm/common/basemodel/cuda_graph.py index 5849cccf54..c080c6f49e 100644 --- a/lightllm/common/basemodel/cuda_graph.py +++ b/lightllm/common/basemodel/cuda_graph.py @@ -9,6 +9,9 @@ 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.common.basemodel.triton_kernel.post_process.vocab_parallel_topk import ( + is_vocab_parallel_topk_enabled, +) from lightllm.utils.torch_memory_saver_utils import ( TorchMemorySaverWrapper, MemoryTag, @@ -279,6 +282,7 @@ def warmup(self, model): b_position_delta=torch.zeros(batch_size, dtype=torch.int32, device="cuda"), is_prefill=False, multimodal_params=[{"images": [], "audios": []} for _ in range(batch_size)], + use_vocab_parallel_topk=is_vocab_parallel_topk_enabled(), **model._gen_special_model_input(batch_size), ) model_output: ModelOutput = model.forward(model_input) @@ -340,6 +344,7 @@ def warmup_overlap(self, model): b_shared_radix_node_id=b_shared_radix_node_id, b_position_delta=torch.zeros(batch_size, dtype=torch.int32, device="cuda"), multimodal_params=[{"images": [], "audios": []} for _ in range(batch_size)], + use_vocab_parallel_topk=is_vocab_parallel_topk_enabled(), **model._gen_special_model_input(batch_size), ) decode_batches.append(micro_batch) diff --git a/lightllm/common/basemodel/infer_struct.py b/lightllm/common/basemodel/infer_struct.py index 29648aa78e..3091c878a9 100755 --- a/lightllm/common/basemodel/infer_struct.py +++ b/lightllm/common/basemodel/infer_struct.py @@ -52,6 +52,8 @@ def __init__(self): self.mem_index: torch.Tensor = None self.return_all_prompt_logics: bool = False + self.use_vocab_parallel_topk: bool = False + self.logits_token_ids: Optional[torch.Tensor] = None # 在开启 return_all_prompt_logics 模式时,保存整个 prefill 阶段每一个 # token 位置的 logits,供后续回传 prompt logprobs 信息使用。 # 仅在 prefill 阶段且需要返回 prompt logprobs 时才会被填充。 diff --git a/lightllm/common/basemodel/prefill_cuda_graph.py b/lightllm/common/basemodel/prefill_cuda_graph.py index bf6039a48f..594dd1ef94 100644 --- a/lightllm/common/basemodel/prefill_cuda_graph.py +++ b/lightllm/common/basemodel/prefill_cuda_graph.py @@ -10,6 +10,9 @@ from lightllm.utils.tensor_utils import tensor_to_no_ref_tensor from lightllm.distributed import dist_group_manager from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.common.basemodel.triton_kernel.post_process.vocab_parallel_topk import ( + is_vocab_parallel_topk_enabled, +) from .infer_struct import InferStateInfo from .cuda_graph import CudaGraph @@ -220,6 +223,7 @@ def warmup(self, model): is_prefill=True, b_prefill_has_output_cpu=[False], multimodal_params=[{"images": [], "audios": []}], + use_vocab_parallel_topk=is_vocab_parallel_topk_enabled(), **model._gen_special_model_input(token_num=total_token_num), ) model_output: ModelOutput = model.forward(model_input) @@ -281,6 +285,7 @@ def warmup_overlap(self, model): is_prefill=True, b_prefill_has_output_cpu=[False], multimodal_params=[{"images": [], "audios": []}], + use_vocab_parallel_topk=is_vocab_parallel_topk_enabled(), **model._gen_special_model_input(token_num=total_token_num), ) diff --git a/lightllm/common/basemodel/triton_kernel/post_process/vocab_parallel_topk.py b/lightllm/common/basemodel/triton_kernel/post_process/vocab_parallel_topk.py new file mode 100644 index 0000000000..bcb48c0090 --- /dev/null +++ b/lightllm/common/basemodel/triton_kernel/post_process/vocab_parallel_topk.py @@ -0,0 +1,94 @@ +"""Collect sparse candidates directly from tensor-parallel vocabulary shards.""" + +import os + +import torch + +from lightllm.distributed.communication_op import all_gather_into_tensor +from lightllm.utils.envs_utils import enable_env_vars + + +VOCAB_PARALLEL_TOPK_ENV = "LIGHTLLM_VOCAB_PARALLEL_TOPK" +VOCAB_PARALLEL_TOPK_SIZE_ENV = "LIGHTLLM_VOCAB_PARALLEL_TOPK_SIZE" +DEFAULT_VOCAB_PARALLEL_TOPK = 128 + + +def is_vocab_parallel_topk_enabled() -> bool: + """Whether target-model greedy batches may use sparse vocabulary output.""" + + return enable_env_vars(VOCAB_PARALLEL_TOPK_ENV) + + +def get_vocab_parallel_topk_size() -> int: + topk = int(os.getenv(VOCAB_PARALLEL_TOPK_SIZE_ENV, str(DEFAULT_VOCAB_PARALLEL_TOPK))) + assert topk > 0, f"{VOCAB_PARALLEL_TOPK_SIZE_ENV} must be positive, got {topk}" + return topk + + +@torch.no_grad() +def vocab_parallel_topk( + local_logits: torch.Tensor, + *, + vocab_size: int, + vocab_start_id: int, + topk: int, + tp_world_size: int, + group, + alloc_func, +) -> tuple[torch.Tensor, torch.Tensor]: + """Gather each TP rank's local top-k logits and their global token ids. + + The returned width is ``tp_world_size * topk``. It intentionally keeps the + union of local candidates: greedy selection remains exact, while probability + calculations over the sparse result are an inexpensive approximation. + """ + + assert local_logits.ndim == 2 and local_logits.is_cuda and local_logits.is_contiguous() + local_vocab_size, token_num = local_logits.shape + # Collectives require every rank to contribute the same shape. Vocabulary + # shards can differ by one row, so cap against the smallest possible shard. + local_topk = min(topk, vocab_size // tp_world_size) + assert local_topk > 0 + assert local_vocab_size >= local_topk + + local_values, local_indexes = torch.topk(local_logits, k=local_topk, dim=0, sorted=False) + local_values = local_values.float() + local_token_ids = local_indexes.to(torch.int32).add_(int(vocab_start_id)) + + if tp_world_size == 1: + candidate_values = local_values.permute(1, 0).contiguous() + candidate_token_ids = local_token_ids.permute(1, 0).contiguous() + else: + # Values and ids are both four bytes. Bit-packing the ids into the FP32 + # payload keeps the operation to one fixed-shape collective. + local_payload = alloc_func( + (local_topk * 2, token_num), + dtype=torch.float32, + device=local_logits.device, + ) + local_payload[:local_topk].copy_(local_values) + local_payload[local_topk:].view(torch.int32).copy_(local_token_ids) + + gathered_payload = alloc_func( + (tp_world_size, local_topk * 2, token_num), + dtype=torch.float32, + device=local_logits.device, + ) + all_gather_into_tensor( + output_=gathered_payload, + input_=local_payload, + group=group, + async_op=False, + ) + candidate_values = gathered_payload[:, :local_topk, :].permute(2, 0, 1).reshape(token_num, -1) + candidate_token_ids = ( + gathered_payload[:, local_topk:, :].view(torch.int32).permute(2, 0, 1).reshape(token_num, -1) + ) + + output_logits = alloc_func( + candidate_values.shape, + dtype=torch.float32, + device=local_logits.device, + ) + output_logits.copy_(candidate_values) + return output_logits, candidate_token_ids.to(torch.int64).contiguous() diff --git a/lightllm/models/gemma4/layer_infer/post_layer_infer.py b/lightllm/models/gemma4/layer_infer/post_layer_infer.py index b736a2d6c1..354d91705d 100644 --- a/lightllm/models/gemma4/layer_infer/post_layer_infer.py +++ b/lightllm/models/gemma4/layer_infer/post_layer_infer.py @@ -5,18 +5,16 @@ class Gemma4PostLayerInfer(LlamaPostLayerInfer): """ Same final RMSNorm + tied lm_head path as Llama, with an extra tanh-based - logit softcap at the end: logits = softcap * tanh(logits / softcap). + transform before sampling: logits = softcap * tanh(logits / softcap). """ def __init__(self, network_config): super().__init__(network_config) self.final_logit_softcapping = float(network_config.get("final_logit_softcapping")) - def token_forward(self, input_embdings, infer_state, layer_weight): - logits = super().token_forward(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 + def _apply_logit_postprocessing(self, logits: torch.Tensor) -> torch.Tensor: + if self.final_logit_softcapping is None or self.final_logit_softcapping <= 0: + return logits + cap = self.final_logit_softcapping + # The historical path materializes FP32 logits before applying softcap. + return torch.tanh(logits.float() / cap) * cap diff --git a/lightllm/models/llama/layer_infer/post_layer_infer.py b/lightllm/models/llama/layer_infer/post_layer_infer.py index 6e4b15a55d..496de58ef1 100644 --- a/lightllm/models/llama/layer_infer/post_layer_infer.py +++ b/lightllm/models/llama/layer_infer/post_layer_infer.py @@ -1,4 +1,3 @@ -import os import torch import torch.functional as F import torch.distributed as dist @@ -7,6 +6,10 @@ from lightllm.models.llama.layer_weights.pre_and_post_layer_weight import LlamaPreAndPostLayerWeight from lightllm.models.llama.infer_struct import LlamaInferStateInfo from lightllm.common.basemodel import PostLayerInferTpl +from lightllm.common.basemodel.triton_kernel.post_process.vocab_parallel_topk import ( + get_vocab_parallel_topk_size, + vocab_parallel_topk, +) from lightllm.distributed.communication_op import all_gather @@ -16,11 +19,17 @@ class LlamaPostLayerInfer(PostLayerInferTpl): def __init__(self, network_config): super().__init__(network_config) self.eps_ = network_config["rms_norm_eps"] + self.vocab_parallel_topk_ = get_vocab_parallel_topk_size() return def _norm(self, input, infer_state, layer_weight: LlamaPreAndPostLayerWeight) -> torch.Tensor: return layer_weight.final_norm_weight_(input=input, eps=self.eps_, alloc_func=self.alloc_tensor) + def _apply_logit_postprocessing(self, logits: torch.Tensor) -> torch.Tensor: + """Apply model-specific transforms while the tensor still contains logits.""" + + return logits + def _slice_get_last_input(self, input_embdings: torch.Tensor, infer_state: LlamaInferStateInfo): embed_dim_ = input_embdings.shape[1] if infer_state.is_prefill: @@ -64,7 +73,7 @@ def _token_forward( 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 + prompt_logics_hiddens, prompt_token_num, layer_weight, infer_state, force_full_logits=True ) return ans_logics @@ -75,13 +84,29 @@ def _lm_head_and_gather( token_num: int, layer_weight: LlamaPreAndPostLayerWeight, infer_state: LlamaInferStateInfo, + force_full_logits: bool = False, ) -> 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 + lm_head = layer_weight.lm_head_weight_ + vocab_size = lm_head.vocab_size + if infer_state.use_vocab_parallel_topk and not force_full_logits: + logic_batch = self._apply_logit_postprocessing(logic_batch) + logits, token_ids = vocab_parallel_topk( + logic_batch, + vocab_size=vocab_size, + vocab_start_id=lm_head.tp_vocab_start_id, + topk=self.vocab_parallel_topk_, + tp_world_size=self.tp_world_size_, + group=infer_state.dist_group, + alloc_func=self.alloc_tensor, + ) + infer_state.logits_token_ids = token_ids + return logits + if self.tp_world_size_ == 1: gather_data = logic_batch else: @@ -98,12 +123,11 @@ def _lm_head_and_gather( 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 + return self._apply_logit_postprocessing(ans_logics) def token_forward( self, input_embdings: torch.Tensor, infer_state: LlamaInferStateInfo, layer_weight: LlamaPreAndPostLayerWeight ): - return self._token_forward(input_embdings=input_embdings, infer_state=infer_state, layer_weight=layer_weight) def overlap_tpsp_token_forward( @@ -114,7 +138,6 @@ def overlap_tpsp_token_forward( infer_state1: LlamaInferStateInfo, layer_weight: BaseLayerWeight, ): - logics = self.token_forward(input_embdings, infer_state, layer_weight=layer_weight) logics1 = self.token_forward(input_embdings1, infer_state1, layer_weight=layer_weight) diff --git a/lightllm/models/qwen3_dspark/layer_infer/post_layer_infer.py b/lightllm/models/qwen3_dspark/layer_infer/post_layer_infer.py index 5a74cd988e..2c89610f13 100644 --- a/lightllm/models/qwen3_dspark/layer_infer/post_layer_infer.py +++ b/lightllm/models/qwen3_dspark/layer_infer/post_layer_infer.py @@ -181,7 +181,12 @@ def token_forward( logits = self._lm_head_and_gather(last_input, token_num, layer_weight, infer_state) block_logits = logits.reshape(num_reqs, self.block_size_, -1) - sampled_tokens = torch.argmax(block_logits, dim=-1) + candidate_indexes = torch.argmax(block_logits, dim=-1) + if infer_state.logits_token_ids is None: + sampled_tokens = candidate_indexes + else: + block_token_ids = infer_state.logits_token_ids.reshape(num_reqs, self.block_size_, -1) + sampled_tokens = block_token_ids.gather(-1, candidate_indexes.unsqueeze(-1)).squeeze(-1) confidence_logits = self.predict_confidence_logits( block_hidden, anchor_token_ids=anchor_token_ids, 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 07b88471f0..0357a25770 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -385,12 +385,19 @@ def _async_copy_next_token_infos_to_pin_mem( ) 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: + def _get_next_token_ranks(self, model_output: ModelOutput, 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 model_output.has_vocab_parallel_logits: + return g_pin_mem_manager.get_const_gpu_tensor( + key="next_token_ranks", + shape=next_token_ids.shape, + fill_value=1 if self.args.enable_rl else -1, + dtype=torch.int32, + ) if not self.args.enable_rl: return g_pin_mem_manager.get_const_gpu_tensor( key="next_token_ranks", @@ -398,8 +405,8 @@ def _get_next_token_ranks(self, logits: torch.Tensor, next_token_ids: torch.Tens 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 + selected_logits = model_output.logits.gather(1, next_token_ids.long().view(-1, 1)) + return (model_output.logits > selected_logits).sum(dim=-1, dtype=torch.int32) + 1 def _capture_prompt_logprobs_if_needed( self, @@ -696,7 +703,6 @@ def _get_classed_reqs( can_alloc_token_num = g_infer_context.get_can_alloc_token_num() for req_obj in ready_reqs: - if req_obj.filter_mark: finished_reqs.append(req_obj) continue @@ -876,17 +882,24 @@ def _trans_req_ids_to_req_objs(self, req_ids: List[int]) -> List[InferReq]: def _gen_argmax_token_ids(self, model_output: ModelOutput): logits = model_output.logits - return torch.argmax(logits, dim=-1) + candidate_indexes = torch.argmax(logits, dim=-1) + return self._map_logits_indexes_to_token_ids(model_output, candidate_indexes) def _gen_argmax_token_ids_and_prob(self, model_output: ModelOutput): logits = model_output.logits - probs = torch.softmax(logits, dim=-1) - max_probs, draft_next_token_ids_gpu = torch.max(probs, dim=-1) - return draft_next_token_ids_gpu, max_probs + max_probs, candidate_indexes = torch.max(torch.softmax(logits, dim=-1), dim=-1) + token_ids = self._map_logits_indexes_to_token_ids(model_output, candidate_indexes) + return token_ids, max_probs + + @staticmethod + def _map_logits_indexes_to_token_ids(model_output: ModelOutput, candidate_indexes: torch.Tensor): + if not model_output.has_vocab_parallel_logits: + return candidate_indexes + return model_output.logits_token_ids.gather(1, candidate_indexes.long().view(-1, 1)).view(-1).long() def _sample_and_scatter_token( self, - logits: torch.Tensor, + model_output: ModelOutput, b_req_idx: torch.Tensor, b_mtp_index: torch.Tensor, run_reqs: List[InferReq], @@ -894,13 +907,14 @@ def _sample_and_scatter_token( b_prefill_has_output_cpu: torch.Tensor = None, mask_func: Optional[Callable] = None, ): - + logits = model_output.logits if mask_func is not None: + assert not model_output.has_vocab_parallel_logits, "constrained sampling requires dense logits" assert len(run_reqs) == logits.shape[0] 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) + next_token_ids, next_token_logprobs = sample(model_output, run_reqs, self.eos_id) + next_token_ranks = self._get_next_token_ranks(model_output, next_token_ids) b_has_out = None if is_prefill: b_has_out = g_pin_mem_manager.gen_from_list( 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 4d09476849..6db9246385 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 @@ -107,11 +107,13 @@ def prefill_normal( ): # 第一阶段: 模型推理 model_input, run_reqs = prepare_prefill_inputs(prefill_reqs, is_chuncked_mode=not self.disable_chunked_prefill) + if self.prefill_mask_func is not None: + model_input.use_vocab_parallel_topk = False 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) (_, next_token_ids_cpu, next_token_logprobs_cpu, next_token_ranks_cpu,) = self._sample_and_scatter_token( - logits=model_output.logits, + model_output=model_output, b_req_idx=model_input.b_req_idx, b_mtp_index=model_input.b_mtp_index, run_reqs=run_reqs, @@ -152,10 +154,12 @@ def decode_normal( decode_reqs: List[InferReq], ): model_input, run_reqs = prepare_decode_inputs(decode_reqs) + if self.decode_mask_func is not None: + model_input.use_vocab_parallel_topk = False 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, next_token_ranks_cpu,) = self._sample_and_scatter_token( - logits=model_output.logits, + model_output=model_output, b_req_idx=model_input.b_req_idx, b_mtp_index=model_input.b_mtp_index, run_reqs=run_reqs, @@ -191,6 +195,8 @@ def prefill_mtp( prefill_reqs: List[InferReq], ): model_input, run_reqs = prepare_prefill_inputs(prefill_reqs, is_chuncked_mode=not self.disable_chunked_prefill) + if self.prefill_mask_func is not None: + model_input.use_vocab_parallel_topk = False 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) @@ -200,7 +206,7 @@ def prefill_mtp( next_token_logprobs_cpu, next_token_ranks_cpu, ) = self._sample_and_scatter_token( - logits=model_output.logits, + model_output=model_output, b_req_idx=model_input.b_req_idx, b_mtp_index=model_input.b_mtp_index, run_reqs=run_reqs, @@ -272,11 +278,11 @@ def decode_mtp( selected_rows = async_selected_row_mask_cpu.tensor.tolist() run_reqs = [req for req, selected in zip(run_reqs, selected_rows) if selected] next_token_ids, next_token_logprobs = sample( - model_output.logits, + model_output, run_reqs, self.eos_id, ) - next_token_ranks = self._get_next_token_ranks(model_output.logits, next_token_ids) + next_token_ranks = self._get_next_token_ranks(model_output, next_token_ids) b_req_mtp_start_loc = gen_b_req_mtp_start_loc(model_input.b_mtp_index, num_reqs=req_num) mtp_accept_len, accepted_index = mtp_utils.verify_mtp_tokens( diff --git a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl_for_reward_model.py b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl_for_reward_model.py index dfb1020820..cabf4f7de0 100644 --- a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl_for_reward_model.py +++ b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl_for_reward_model.py @@ -14,9 +14,9 @@ def __init__(self) -> None: return def reward_prefill(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq]): - 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_input.use_vocab_parallel_topk = False model_output = self.model.forward(model_input) scores: torch.Tensor = model_output.logits 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 21979cbef0..204e1312e3 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 @@ -40,10 +40,7 @@ def beam_prefill(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq ) with torch.cuda.stream(g_infer_context.get_overlap_stream()): - model_output = self.model.forward(model_input) - logits = model_output.logits - batch_idx, run_reqs = self._diverse_copy( master_reqs=group_reqs, b_prefill_has_out=model_input.b_prefill_has_output_cpu ) @@ -60,11 +57,11 @@ def beam_prefill(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq non_blocking=True ) - logits = logits[batch_idx] + sampled_output = model_output.index_select_logits_rows(batch_idx) 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) + next_token_ids, next_token_logprobs = sample(sampled_output, run_reqs, self.eos_id) + next_token_ranks = self._get_next_token_ranks(sampled_output, next_token_ids) scatter_token( next_token_ids=next_token_ids, 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 9a81927bc1..fb5aea673a 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 @@ -187,7 +187,7 @@ def prefill_normal( next_token_logprobs_cpu, next_token_ranks_cpu, ) = self._sample_and_scatter_token( - logits=model_output.logits, + model_output=model_output, b_req_idx=model_input.b_req_idx, b_mtp_index=model_input.b_mtp_index, run_reqs=run_reqs, @@ -240,7 +240,7 @@ def decode_normal(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq next_token_logprobs_cpu, next_token_ranks_cpu, ) = self._sample_and_scatter_token( - logits=model_output.logits, + model_output=model_output, b_req_idx=model_input.b_req_idx, b_mtp_index=model_input.b_mtp_index, run_reqs=run_reqs, @@ -287,13 +287,8 @@ def prefill_overlap(self, event_pack: OverlapEventPack, prefill_reqs: List[Infer 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) - logits = torch.empty((req_num0 + req_num1, logits0.shape[1]), dtype=logits0.dtype, device=logits0.device) - logits[0:req_num0, :].copy_(logits0, non_blocking=True) - logits[req_num0 : req_num0 + req_num1, :].copy_(logits1, non_blocking=True) + sampled_output = ModelOutput.concat_logits_rows([model_output0, model_output1]) run_reqs = run_reqs0 + run_reqs1 b_has_out_cpu = model_input0.b_prefill_has_output_cpu + model_input1.b_prefill_has_output_cpu @@ -307,7 +302,7 @@ def prefill_overlap(self, event_pack: OverlapEventPack, prefill_reqs: List[Infer next_token_logprobs_cpu, next_token_ranks_cpu, ) = self._sample_and_scatter_token( - logits=logits, + model_output=sampled_output, b_req_idx=b_req_idx, b_mtp_index=b_mtp_index, run_reqs=run_reqs, @@ -356,7 +351,7 @@ def decode_overlap(self, event_pack: OverlapEventPack, decode_reqs: List[InferRe with torch.cuda.stream(g_infer_context.get_overlap_stream()): model_output0, model_output1 = self.model.microbatch_overlap_decode(model_input0, model_input1) if req_num0 + req_num1 > 0: - logits = torch.cat((model_output0.logits, model_output1.logits), dim=0) + sampled_output = ModelOutput.concat_logits_rows([model_output0, model_output1]) b_req_idx = torch.cat((model_input0.b_req_idx, model_input1.b_req_idx), dim=0) b_mtp_index = torch.cat((model_input0.b_mtp_index, model_input1.b_mtp_index), dim=0) ( @@ -365,7 +360,7 @@ def decode_overlap(self, event_pack: OverlapEventPack, decode_reqs: List[InferRe next_token_logprobs_cpu, next_token_ranks_cpu, ) = self._sample_and_scatter_token( - logits=logits, + model_output=sampled_output, b_req_idx=b_req_idx, b_mtp_index=b_mtp_index, run_reqs=run_reqs, @@ -421,7 +416,7 @@ def prefill_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq] next_token_logprobs_cpu, next_token_ranks_cpu, ) = self._sample_and_scatter_token( - logits=model_output.logits, + model_output=model_output, b_req_idx=b_req_idx, b_mtp_index=b_mtp_index, run_reqs=run_reqs, @@ -446,7 +441,6 @@ def prefill_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq] sync_event.record() if req_num > 0: - # 第二阶段 event_pack.notify_post_handle_and_wait_pre_post_handle() update_packs = self._pre_post_handle(run_reqs, is_chuncked_mode=not self.disable_chunked_prefill) @@ -499,11 +493,11 @@ def decode_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): if req_num > 0: next_token_ids, next_token_logprobs = sample( - model_output.logits, + model_output, run_reqs, self.eos_id, ) - next_token_ranks = self._get_next_token_ranks(model_output.logits, next_token_ids) + next_token_ranks = self._get_next_token_ranks(model_output, next_token_ids) b_req_mtp_start_loc = gen_b_req_mtp_start_loc( b_mtp_index=model_input.b_mtp_index, @@ -649,17 +643,9 @@ def prefill_overlap_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[I 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) req_num = req_num0 + req_num1 - logits = torch.empty( - (req_num0 + req_num1, logits0.shape[1]), - dtype=logits0.dtype, - device=logits0.device, - ) - logits[0:req_num0, :].copy_(logits0, non_blocking=True) - logits[req_num0 : (req_num0 + req_num1), :].copy_(logits1, non_blocking=True) + sampled_output = ModelOutput.concat_logits_rows([model_output0, model_output1]) run_reqs = run_reqs0 + run_reqs1 b_has_out_cpu = model_input0.b_prefill_has_output_cpu + model_input1.b_prefill_has_output_cpu @@ -673,7 +659,7 @@ def prefill_overlap_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[I next_token_logprobs_cpu, next_token_ranks_cpu, ) = self._sample_and_scatter_token( - logits=logits, + model_output=sampled_output, run_reqs=run_reqs, b_req_idx=b_req_idx, b_mtp_index=b_mtp_index, @@ -681,7 +667,7 @@ def prefill_overlap_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[I b_prefill_has_output_cpu=b_has_out_cpu, ) else: - next_token_ids = torch.empty((0,), dtype=torch.int64, device=logits.device) + next_token_ids = torch.empty((0,), dtype=torch.int64, device=sampled_output.logits.device) target_next_token_ids_gpu0 = next_token_ids[:req_num0] target_next_token_ids_gpu1 = next_token_ids[req_num0:] @@ -769,20 +755,12 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf verify_row_num0 = model_input0.batch_size verify_row_num1 = model_input1.batch_size verify_row_num = verify_row_num0 + verify_row_num1 - logits0 = model_output0.logits - logits1 = model_output1.logits run_reqs = run_reqs0 + run_reqs1 if req_num > 0: assert len(run_reqs) == verify_row_num - logits = torch.empty( - (verify_row_num, logits0.shape[1]), - dtype=logits0.dtype, - device=logits0.device, - ) - logits[:verify_row_num0, :].copy_(logits0, non_blocking=True) - logits[verify_row_num0:, :].copy_(logits1, non_blocking=True) - 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) + sampled_output = ModelOutput.concat_logits_rows([model_output0, model_output1]) + next_token_ids, next_token_logprobs = sample(sampled_output, run_reqs, self.eos_id) + next_token_ranks = self._get_next_token_ranks(sampled_output, next_token_ids) ( next_token_ids_cpu, next_token_logprobs_cpu, diff --git a/lightllm/server/router/model_infer/mode_backend/generic_post_process.py b/lightllm/server/router/model_infer/mode_backend/generic_post_process.py index 5b29ea0510..bd20f50d1f 100644 --- a/lightllm/server/router/model_infer/mode_backend/generic_post_process.py +++ b/lightllm/server/router/model_infer/mode_backend/generic_post_process.py @@ -1,5 +1,9 @@ import torch from typing import List, Tuple +from lightllm.common.basemodel.batch_objs import ModelOutput +from lightllm.common.basemodel.triton_kernel.post_process.vocab_parallel_topk import ( + is_vocab_parallel_topk_enabled, +) from lightllm.common.basemodel.triton_kernel.post_process.apply_penalty import apply_penalty from lightllm.common.basemodel.triton_kernel.post_process.apply_penalty_gpu_cache import apply_penalty_gpu_cache from lightllm.common.basemodel.triton_kernel.post_process.apply_invalid_token import apply_invalid_token_ids @@ -8,7 +12,44 @@ from lightllm.utils.envs_utils import get_env_start_args -def sample(logits: torch.Tensor, reqs: List[InferReq], eos_id: List[int] = [2]): +def _can_use_unmodified_greedy_logits(reqs: List[InferReq]) -> bool: + """Whether sampling is exactly argmax over the incoming logits.""" + + for req_obj in reqs: + sample_param = req_obj.sampling_param + shm_param = sample_param.shm_param + if shm_param.top_k != 1 or shm_param.temperature != 1.0: + return False + if ( + shm_param.presence_penalty != 0.0 + or shm_param.frequency_penalty != 0.0 + or shm_param.repetition_penalty != 1.0 + ): + return False + if shm_param.exponential_decay_length_penalty.to_tuple()[1] != 1.0: + return False + out_token_len = req_obj.get_cur_total_len() - req_obj.shm_req.input_len + if out_token_len < shm_param.min_new_tokens - 1: + return False + if sample_param.invalid_token_ids: + return False + return True + + +def can_use_vocab_parallel_topk(reqs: List[InferReq]) -> bool: + return is_vocab_parallel_topk_enabled() and _can_use_unmodified_greedy_logits(reqs) + + +def sample(model_output: ModelOutput, reqs: List[InferReq], eos_id: List[int] = [2]): + logits = model_output.logits + if model_output.has_vocab_parallel_logits: + if not _can_use_unmodified_greedy_logits(reqs): + raise RuntimeError("vocab-parallel top-k logits require unmodified greedy requests") + probs = torch.softmax(logits, dim=-1) + max_probs, candidate_indexes = torch.max(probs, dim=-1) + token_ids = model_output.logits_token_ids.gather(1, candidate_indexes.view(-1, 1)).view(-1).long() + return token_ids, torch.log(max_probs) + ( b_req_idx, b_temperatures, diff --git a/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py b/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py index 22731439c4..ee95ff24c5 100644 --- a/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py +++ b/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py @@ -3,6 +3,9 @@ from typing import List, Tuple from lightllm.server.router.model_infer.infer_batch import InferReq, g_infer_context from lightllm.common.basemodel.batch_objs import ModelInput +from lightllm.server.router.model_infer.mode_backend.generic_post_process import ( + can_use_vocab_parallel_topk, +) INT64_MAX = torch.iinfo(torch.int64).max @@ -87,6 +90,7 @@ def prepare_prefill_inputs(req_objs: List[InferReq], is_chuncked_mode: bool) -> is_prefill=True, b_prefill_has_output_cpu=b_prefill_has_output, multimodal_params=batch_multimodal_params, + use_vocab_parallel_topk=can_use_vocab_parallel_topk(run_reqs), ) return model_input, run_reqs @@ -160,6 +164,7 @@ def prepare_decode_inputs(req_objs: List[InferReq]) -> Tuple[ModelInput, List[In b_shared_radix_node_id=b_shared_radix_node_id, is_prefill=False, multimodal_params=multimodal_params, + use_vocab_parallel_topk=can_use_vocab_parallel_topk(run_reqs), ) return model_input, run_reqs @@ -176,6 +181,9 @@ def overlap_prepare_decode_inputs(req_objs: List[InferReq]): model_input1, run_reqs1 = prepare_decode_inputs( req_objs=decode_reqs1, ) + use_vocab_parallel_topk = can_use_vocab_parallel_topk(run_reqs0 + run_reqs1) + model_input0.use_vocab_parallel_topk = use_vocab_parallel_topk + model_input1.use_vocab_parallel_topk = use_vocab_parallel_topk return model_input0, run_reqs0, decode_reqs0, model_input1, run_reqs1, decode_reqs1 @@ -211,6 +219,9 @@ def overlap_prepare_prefill_inputs(req_objs: List[InferReq]): req_objs=right_reqs, is_chuncked_mode=True, ) + use_vocab_parallel_topk = can_use_vocab_parallel_topk(run_reqs0 + run_reqs1) + model_input0.use_vocab_parallel_topk = use_vocab_parallel_topk + model_input1.use_vocab_parallel_topk = use_vocab_parallel_topk return model_input0, run_reqs0, model_input1, run_reqs1 diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py index 6b8c23e8fd..137a580ae6 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py @@ -141,7 +141,7 @@ def propose_next_overlap( req_num_by_batch, ) ): - accepted_tail_output = ModelOutput(logits=extend_output.logits.index_select(0, accepted_tail_rows)) + accepted_tail_output = extend_output.index_select_logits_rows(accepted_tail_rows) if self.enable_dynmaic_mtp: draft_token_ids, draft_token_probs = self._gen_argmax_token_ids_and_prob(accepted_tail_output) draft_token_probs = draft_token_probs.float() diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py index 3d2c0a0e86..f2c730ccd5 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py @@ -81,7 +81,7 @@ def propose_next( # 只在 req_num 行 logits 上进行 argmax,避免为未接受的 verify 行执行 # vocabulary reduction。第一列 proposal 来自每个请求的 accepted tail。 - accepted_tail_output = ModelOutput(logits=extend_output.logits.index_select(0, accepted_tail_rows)) + accepted_tail_output = extend_output.index_select_logits_rows(accepted_tail_rows) if self.enable_dynmaic_mtp: draft_token_ids, draft_token_probs = self._gen_argmax_token_ids_and_prob(accepted_tail_output) schedule_scores_by_step.append(draft_token_probs.float().unsqueeze(1)) diff --git a/unit_tests/common/basemodel/test_model_output.py b/unit_tests/common/basemodel/test_model_output.py index 6f9477e294..68a5286774 100644 --- a/unit_tests/common/basemodel/test_model_output.py +++ b/unit_tests/common/basemodel/test_model_output.py @@ -7,20 +7,36 @@ from lightllm.common.basemodel.batch_objs import ModelInput, ModelMtpOutputCollector, ModelOutput +def test_vocab_parallel_metadata_follows_row_operations(): + output = ModelOutput( + logits=torch.tensor([[8.0, 1.0], [7.0, 2.0], [6.0, 3.0]]), + logits_token_ids=torch.tensor([[4, 14], [9, 19], [2, 12]]), + ) + + selected = output.index_select_logits_rows(torch.tensor([2, 0])) + combined = ModelOutput.concat_logits_rows([selected, output.index_select_logits_rows(torch.tensor([1]))]) + + torch.testing.assert_close(combined.logits, torch.tensor([[6.0, 3.0], [8.0, 1.0], [7.0, 2.0]])) + torch.testing.assert_close(combined.logits_token_ids, torch.tensor([[2, 12], [4, 14], [9, 19]])) + + def test_decode_unpad_slices_spec_output_with_logits(): model = TpPartBaseModel.__new__(TpPartBaseModel) output = ModelOutput( logits=torch.arange(24).view(6, 4), + logits_token_ids=torch.arange(100, 124).view(6, 4), mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.arange(18).view(6, 3)), ) unpadded = model._create_unpad_decode_model_output(output, origin_batch_size=4) assert unpadded.logits.shape == (4, 4) + assert unpadded.logits_token_ids.shape == (4, 4) assert unpadded.mtp_collector.spec_hidden.shape == (4, 3) # Unpadding returns a shallow output copy and leaves the graph-owned # tensors on the original ModelOutput intact. assert output.logits.shape == (6, 4) + assert output.logits_token_ids.shape == (6, 4) assert output.mtp_collector.spec_hidden.shape == (6, 3) @@ -28,6 +44,7 @@ def test_prefill_unpad_uses_token_rows_for_spec_hidden(): model = TpPartBaseModel.__new__(TpPartBaseModel) output = ModelOutput( logits=torch.arange(20).view(5, 4), + logits_token_ids=torch.arange(100, 120).view(5, 4), mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.arange(24).view(8, 3)), prompt_logics=torch.arange(32).view(8, 4), ) @@ -39,6 +56,7 @@ def test_prefill_unpad_uses_token_rows_for_spec_hidden(): ) assert unpadded.logits.shape == (3, 4) + assert unpadded.logits_token_ids.shape == (3, 4) assert unpadded.mtp_collector.spec_hidden.shape == (6, 3) assert unpadded.prompt_logics.shape == (6, 4) @@ -94,6 +112,48 @@ def _create_empty_decode_input(): ) +def test_infer_state_enables_vocab_parallel_topk_for_draft_or_requested_target(monkeypatch): + model = TpPartBaseModel.__new__(TpPartBaseModel) + model.infer_state_class = basemodel.InferStateInfo + model.hidden_collector_prototype = SimpleNamespace(new_instance=lambda: object()) + model.is_token_healing = False + model.return_all_prompt_logics = False + model.is_mtp_draft_model = False + model.mem_manager = object() + model.req_manager = object() + model.decode_att_backend = SimpleNamespace(create_att_decode_state=lambda infer_state: None) + model.decode_att_backend1 = None + monkeypatch.setattr(basemodel.dist_group_manager, "get_group", lambda _: None) + + model_input = _create_empty_decode_input() + infer_state = model._create_inferstate(model_input) + assert not infer_state.use_vocab_parallel_topk + + model_input.use_vocab_parallel_topk = True + infer_state = model._create_inferstate(model_input) + assert infer_state.use_vocab_parallel_topk + + model.is_mtp_draft_model = True + model_input.use_vocab_parallel_topk = False + infer_state = model._create_inferstate(model_input) + assert infer_state.use_vocab_parallel_topk + + +def test_cuda_graph_contract_falls_back_for_dense_target_batch(monkeypatch): + model = TpPartBaseModel.__new__(TpPartBaseModel) + model.is_mtp_draft_model = False + sparse_input = SimpleNamespace(use_vocab_parallel_topk=True) + dense_input = SimpleNamespace(use_vocab_parallel_topk=False) + + monkeypatch.setattr(basemodel, "is_vocab_parallel_topk_enabled", lambda: True) + assert model._is_cuda_graph_output_compatible(sparse_input) + assert not model._is_cuda_graph_output_compatible(dense_input) + assert not model._is_cuda_graph_output_compatible(sparse_input, dense_input) + + model.is_mtp_draft_model = True + assert model._is_cuda_graph_output_compatible(dense_input) + + @torch.no_grad() def test_decode_pads_only_once_after_selecting_execution_path(monkeypatch): monkeypatch.setattr(basemodel, "copy_kv_index_to_req", lambda *args: None) @@ -117,6 +177,7 @@ def test_decode_pads_only_once_after_selecting_execution_path(monkeypatch): ) in execution_configs: model = TpPartBaseModel.__new__(TpPartBaseModel) model.args = SimpleNamespace(enable_tpsp_mix_mode=enable_tpsp_mix_mode) + model.is_mtp_draft_model = False model.tp_world_size_ = tp_world_size model.mem_manager = SimpleNamespace(HOLD_TOKEN_MEMINDEX=99) model.req_manager = SimpleNamespace(HOLD_REQUEST_ID=88, req_to_token_indexs=object()) diff --git a/unit_tests/common/basemodel/triton_kernel/test_vocab_parallel_topk.py b/unit_tests/common/basemodel/triton_kernel/test_vocab_parallel_topk.py new file mode 100644 index 0000000000..1b7c2b4dc9 --- /dev/null +++ b/unit_tests/common/basemodel/triton_kernel/test_vocab_parallel_topk.py @@ -0,0 +1,81 @@ +import importlib + +import pytest +import torch + + +pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") + + +@pytest.mark.parametrize("token_num", [1, 7, 64]) +def test_vocab_parallel_topk_collects_candidates_and_preserves_global_argmax(monkeypatch, token_num): + module = importlib.import_module("lightllm.common.basemodel.triton_kernel.post_process.vocab_parallel_topk") + tp_world_size = 4 + local_vocab_size = 1024 + vocab_size = tp_world_size * local_vocab_size + local_topk = 16 + generator = torch.Generator(device="cuda").manual_seed(20260902 + token_num) + local_logits_by_rank = [ + torch.randn( + (local_vocab_size, token_num), + dtype=torch.bfloat16, + device="cuda", + generator=generator, + ) + for _ in range(tp_world_size) + ] + local_logits_by_rank[3][2, 0] = 20.0 + + def pack_rank(rank): + values, indexes = torch.topk(local_logits_by_rank[rank], k=local_topk, dim=0, sorted=False) + payload = torch.empty((local_topk * 2, token_num), dtype=torch.float32, device="cuda") + payload[:local_topk].copy_(values.float()) + payload[local_topk:].view(torch.int32).copy_(indexes.to(torch.int32) + rank * local_vocab_size) + return payload + + payloads = [pack_rank(rank) for rank in range(tp_world_size)] + + def fake_all_gather_into_tensor(output_, input_, **_kwargs): + assert output_.shape == (tp_world_size, local_topk * 2, token_num) + assert input_.shape == (local_topk * 2, token_num) + for output, payload in zip(output_, payloads): + output.copy_(payload) + + monkeypatch.setattr(module, "all_gather_into_tensor", fake_all_gather_into_tensor) + actual_logits, actual_ids = module.vocab_parallel_topk( + local_logits_by_rank[0], + vocab_size=vocab_size, + vocab_start_id=0, + topk=local_topk, + tp_world_size=tp_world_size, + group=None, + alloc_func=torch.empty, + ) + + assert actual_logits.shape == (token_num, tp_world_size * local_topk) + assert actual_ids.shape == actual_logits.shape + assert actual_logits.dtype == torch.float32 + assert actual_ids.dtype == torch.int64 + + full_logits = torch.cat(local_logits_by_rank, dim=0).transpose(0, 1).float() + candidate_indexes = actual_logits.argmax(dim=1, keepdim=True) + actual_argmax_ids = actual_ids.gather(1, candidate_indexes).view(-1) + torch.testing.assert_close(actual_argmax_ids, full_logits.argmax(dim=1), rtol=0, atol=0) + + +def test_vocab_parallel_topk_single_rank_caps_width_to_vocabulary(): + module = importlib.import_module("lightllm.common.basemodel.triton_kernel.post_process.vocab_parallel_topk") + local_logits = torch.tensor([[1.0], [3.0], [2.0]], device="cuda") + + logits, token_ids = module.vocab_parallel_topk( + local_logits, + vocab_size=3, + vocab_start_id=0, + topk=128, + tp_world_size=1, + group=None, + alloc_func=torch.empty, + ) + + assert logits.shape == (1, 3) + torch.testing.assert_close(token_ids.sort(dim=1).values, torch.tensor([[0, 1, 2]], device="cuda")) diff --git a/unit_tests/models/test_gemma4_vocab_parallel_topk.py b/unit_tests/models/test_gemma4_vocab_parallel_topk.py new file mode 100644 index 0000000000..13ff5c3c24 --- /dev/null +++ b/unit_tests/models/test_gemma4_vocab_parallel_topk.py @@ -0,0 +1,67 @@ +from types import SimpleNamespace + +import torch + +import lightllm.models.llama.layer_infer.post_layer_infer as llama_post_layer +from lightllm.models.gemma4.layer_infer.post_layer_infer import Gemma4PostLayerInfer + + +def test_vocab_parallel_topk_softcaps_local_logits_before_candidate_selection(monkeypatch): + post = Gemma4PostLayerInfer.__new__(Gemma4PostLayerInfer) + post.final_logit_softcapping = 2.0 + post.vocab_parallel_topk_ = 2 + post.tp_world_size_ = 1 + post.alloc_tensor = torch.empty + post._norm = lambda hidden, infer_state, layer_weight: hidden + + local_logits = torch.tensor( + [[4.0, -4.0], [2.0, -2.0], [1.0, -1.0]], + dtype=torch.bfloat16, + ) + + class LMHead: + vocab_size = 3 + tp_vocab_start_id = 0 + + def __call__(self, input, alloc_func): + return local_logits + + sparse_logits = torch.tensor([[1.5, 1.0], [-1.0, -1.5]]) + token_ids = torch.tensor([[0, 1], [2, 1]]) + captured = {} + + def fake_vocab_parallel_topk(logits, **kwargs): + captured["logits"] = logits + return sparse_logits, token_ids + + monkeypatch.setattr(llama_post_layer, "vocab_parallel_topk", fake_vocab_parallel_topk) + infer_state = SimpleNamespace( + dist_group=None, + use_vocab_parallel_topk=True, + logits_token_ids=None, + ) + + result = post._lm_head_and_gather( + hidden=torch.empty((2, 3)), + token_num=2, + layer_weight=SimpleNamespace(lm_head_weight_=LMHead()), + infer_state=infer_state, + ) + + expected = torch.tanh(local_logits.float() / 2.0) * 2.0 + assert captured["logits"].dtype == torch.float32 + torch.testing.assert_close(captured["logits"], expected) + assert result is sparse_logits + assert infer_state.logits_token_ids is token_ids + + +def test_full_logits_softcap_after_float32_conversion(): + post = Gemma4PostLayerInfer.__new__(Gemma4PostLayerInfer) + post.final_logit_softcapping = 2.0 + logits = torch.tensor([[1.234375, -3.140625]], dtype=torch.bfloat16) + + actual = post._apply_logit_postprocessing(logits) + expected = torch.tanh(logits.float() / 2.0) * 2.0 + + assert actual.dtype == torch.float32 + torch.testing.assert_close(actual, expected) diff --git a/unit_tests/models/test_vocab_parallel_topk_output.py b/unit_tests/models/test_vocab_parallel_topk_output.py new file mode 100644 index 0000000000..e8dc0adb2d --- /dev/null +++ b/unit_tests/models/test_vocab_parallel_topk_output.py @@ -0,0 +1,72 @@ +from types import SimpleNamespace + +import torch + +from lightllm.common.basemodel.batch_objs import ModelOutput +from lightllm.models.qwen3_dspark.layer_infer.post_layer_infer import Qwen3DSparkPostLayerInfer +from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend + + +def test_argmax_restores_global_token_ids_and_uses_candidate_probability(): + backend = ModeBackend.__new__(ModeBackend) + output = ModelOutput( + logits=torch.tensor([[3.0, 1.0, 2.0], [0.0, 5.0, 4.0]]), + logits_token_ids=torch.tensor([[30, 10, 20], [100, 500, 400]]), + ) + + token_ids = backend._gen_argmax_token_ids(output) + token_ids_with_prob, probs = backend._gen_argmax_token_ids_and_prob(output) + + torch.testing.assert_close(token_ids, torch.tensor([30, 500])) + torch.testing.assert_close(token_ids_with_prob, token_ids) + torch.testing.assert_close(probs, torch.softmax(output.logits, dim=-1).max(dim=-1).values) + + +def test_dense_argmax_keeps_column_index_semantics(): + backend = ModeBackend.__new__(ModeBackend) + output = ModelOutput(logits=torch.tensor([[1.0, 4.0, 2.0]])) + + torch.testing.assert_close(backend._gen_argmax_token_ids(output), torch.tensor([1])) + + +def test_dspark_confidence_path_receives_mapped_sparse_token_ids(): + post = Qwen3DSparkPostLayerInfer.__new__(Qwen3DSparkPostLayerInfer) + post.block_size_ = 2 + post.markov_rank_ = 0 + post._slice_get_last_input = lambda input_embeddings, infer_state: (input_embeddings, 4) + sparse_logits = torch.tensor([[1.0, 4.0], [5.0, 2.0], [3.0, 7.0], [9.0, 8.0]]) + sparse_token_ids = torch.tensor([[10, 40], [50, 20], [30, 70], [90, 80]]) + + def gather_vocab_parallel(*args, **kwargs): + infer_state = args[3] + infer_state.logits_token_ids = sparse_token_ids + return sparse_logits + + post._lm_head_and_gather = gather_vocab_parallel + observed = {} + + def predict_confidence(block_hidden, anchor_token_ids, sampled_tokens, layer_weight): + observed["sampled_tokens"] = sampled_tokens + return None + + post.predict_confidence_logits = predict_confidence + + class Collector: + def add_mtp_outputs(self, **kwargs): + self.outputs = kwargs + + infer_state = SimpleNamespace( + is_prefill=False, + input_ids=torch.tensor([1, 0, 2, 0]), + logits_token_ids=None, + hidden_collector=Collector(), + ) + + returned_logits = post.token_forward( + input_embdings=torch.ones((4, 3)), + infer_state=infer_state, + layer_weight=object(), + ) + + torch.testing.assert_close(returned_logits, sparse_logits) + torch.testing.assert_close(observed["sampled_tokens"], torch.tensor([[40, 50], [70, 90]])) diff --git a/unit_tests/server/router/model_infer/mode_backend/test_vocab_parallel_topk_sampling.py b/unit_tests/server/router/model_infer/mode_backend/test_vocab_parallel_topk_sampling.py new file mode 100644 index 0000000000..09f74e3b07 --- /dev/null +++ b/unit_tests/server/router/model_infer/mode_backend/test_vocab_parallel_topk_sampling.py @@ -0,0 +1,91 @@ +from types import SimpleNamespace + +import pytest +import torch + +from lightllm.common.basemodel.batch_objs import ModelOutput +from lightllm.server.router.model_infer.mode_backend.generic_post_process import ( + _can_use_unmodified_greedy_logits, + can_use_vocab_parallel_topk, + sample, +) +from lightllm.utils.envs_utils import enable_env_vars + + +def make_req(**overrides): + values = { + "top_k": 1, + "temperature": 1.0, + "presence_penalty": 0.0, + "frequency_penalty": 0.0, + "repetition_penalty": 1.0, + "min_new_tokens": 1, + "decay_factor": 1.0, + "invalid_token_ids": [], + "output_len": 0, + } + values.update(overrides) + shm_param = SimpleNamespace( + top_k=values["top_k"], + temperature=values["temperature"], + presence_penalty=values["presence_penalty"], + frequency_penalty=values["frequency_penalty"], + repetition_penalty=values["repetition_penalty"], + min_new_tokens=values["min_new_tokens"], + exponential_decay_length_penalty=SimpleNamespace(to_tuple=lambda: (1, values["decay_factor"])), + ) + input_len = 10 + return SimpleNamespace( + sampling_param=SimpleNamespace( + shm_param=shm_param, + invalid_token_ids=values["invalid_token_ids"], + ), + shm_req=SimpleNamespace(input_len=input_len), + get_cur_total_len=lambda: input_len + values["output_len"], + ) + + +def test_accepts_unmodified_greedy_requests(): + assert _can_use_unmodified_greedy_logits([make_req(), make_req(output_len=5)]) + + +def test_feature_gate_requires_environment_and_eligible_batch(monkeypatch): + monkeypatch.delenv("LIGHTLLM_VOCAB_PARALLEL_TOPK", raising=False) + enable_env_vars.cache_clear() + assert not can_use_vocab_parallel_topk([make_req()]) + + monkeypatch.setenv("LIGHTLLM_VOCAB_PARALLEL_TOPK", "1") + enable_env_vars.cache_clear() + assert can_use_vocab_parallel_topk([make_req()]) + assert not can_use_vocab_parallel_topk([make_req(top_k=2)]) + enable_env_vars.cache_clear() + + +def test_samples_sparse_candidates_and_maps_global_token_ids(): + model_output = ModelOutput( + logits=torch.tensor([[9.0, 7.0, 1.0], [4.0, 6.0, 5.0]], dtype=torch.float32), + logits_token_ids=torch.tensor([[17, 3, 8], [3, 20, 11]], dtype=torch.int64), + ) + + token_ids, token_logprobs = sample(model_output, [make_req(), make_req()]) + + expected_probs = torch.softmax(model_output.logits, dim=-1).max(dim=-1).values + torch.testing.assert_close(token_ids, torch.tensor([17, 20])) + torch.testing.assert_close(token_logprobs, torch.log(expected_probs)) + + +@pytest.mark.parametrize( + "override", + [ + {"top_k": 2}, + {"temperature": 0.5}, + {"presence_penalty": 0.1}, + {"frequency_penalty": 0.1}, + {"repetition_penalty": 1.1}, + {"decay_factor": 1.1}, + {"invalid_token_ids": [7]}, + {"min_new_tokens": 2}, + ], +) +def test_rejects_logits_modifiers(override): + assert not _can_use_unmodified_greedy_logits([make_req(**override)])