Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
19 changes: 19 additions & 0 deletions lightllm/common/basemodel/basemodel.py
Original file line number Diff line number Diff line change
Expand Up @@ -384,6 +384,7 @@ def _create_inferstate(self, model_input: ModelInput, microbatch_index: int = 0)
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_greedy = self.is_mtp_draft_model
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
Expand Down Expand Up @@ -534,6 +535,9 @@ 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.logits_logsumexp = new_model_output.logits_logsumexp[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,
Expand All @@ -546,6 +550,9 @@ 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.logits_logsumexp = new_model_output.logits_logsumexp[0:origin_batch_size]
new_model_output.mtp_collector = padded_model_output.mtp_collector.unpad_prefill(
origin_handle_token_num=origin_handle_token_num
)
Expand Down Expand Up @@ -737,6 +744,8 @@ 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,
logits_logsumexp=infer_state.logits_logsumexp,
mtp_collector=infer_state.hidden_collector.finish_output(infer_state=infer_state),
prompt_logics=infer_state.prompt_logics,
)
Expand Down Expand Up @@ -766,6 +775,8 @@ 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,
logits_logsumexp=infer_state.logits_logsumexp,
mtp_collector=infer_state.hidden_collector.finish_output(infer_state=infer_state),
)

Expand Down Expand Up @@ -1020,11 +1031,15 @@ 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,
logits_logsumexp=infer_state.logits_logsumexp,
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,
logits_logsumexp=infer_state1.logits_logsumexp,
mtp_collector=infer_state1.hidden_collector.finish_output(infer_state=infer_state1),
prompt_logics=infer_state1.prompt_logics,
)
Expand Down Expand Up @@ -1069,10 +1084,14 @@ 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,
logits_logsumexp=infer_state.logits_logsumexp,
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,
logits_logsumexp=infer_state1.logits_logsumexp,
mtp_collector=infer_state1.hidden_collector.finish_output(infer_state=infer_state1),
)

Expand Down
36 changes: 36 additions & 0 deletions lightllm/common/basemodel/batch_objs.py
Original file line number Diff line number Diff line change
Expand Up @@ -200,10 +200,46 @@ class ModelOutput:
# 需要返回 prompt logprobs 信息时才会非空。
prompt_logics: Optional[torch.Tensor] = None

# Vocab-parallel outputs keep logits as logits while mapping each sparse
# column back to its global token id. logits_logsumexp is computed over the
# complete vocabulary, so sparse argmax probabilities remain exact.
# Both fields are None for historical dense logits.
logits_token_ids: Optional[torch.Tensor] = None
logits_logsumexp: Optional[torch.Tensor] = None

def __post_init__(self) -> None:
if self.mtp_collector is None:
self.mtp_collector = ModelMtpOutputCollector()
assert (self.logits_token_ids is None) == (self.logits_logsumexp is None)
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
assert self.logits_logsumexp.shape == (self.logits.shape[0],)
assert self.logits_logsumexp.dtype == torch.float32
assert self.logits_logsumexp.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.logits_logsumexp = tensor_to_no_ref_tensor(self.logits_logsumexp)
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
),
logits_logsumexp=(
self.logits_logsumexp.index_select(0, index) if self.logits_logsumexp is not None else None
),
)
3 changes: 3 additions & 0 deletions lightllm/common/basemodel/infer_struct.py
Original file line number Diff line number Diff line change
Expand Up @@ -52,6 +52,9 @@ def __init__(self):
self.mem_index: torch.Tensor = None

self.return_all_prompt_logics: bool = False
self.use_vocab_parallel_greedy: bool = False
self.logits_token_ids: Optional[torch.Tensor] = None
self.logits_logsumexp: Optional[torch.Tensor] = None
# 在开启 return_all_prompt_logics 模式时,保存整个 prefill 阶段每一个
# token 位置的 logits,供后续回传 prompt logprobs 信息使用。
# 仅在 prefill 阶段且需要返回 prompt logprobs 时才会被填充。
Expand Down
109 changes: 109 additions & 0 deletions lightllm/common/basemodel/triton_kernel/post_process/greedy_sample.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,109 @@
"""Local greedy statistics for distributed vocabulary shards."""

import torch
import triton
import triton.language as tl


@triton.jit
def _greedy_sample_stage1_kernel(
logits,
partial_max,
partial_sum,
partial_argmax,
stride_row,
stride_col,
vocab_size: tl.constexpr,
num_blocks: tl.constexpr,
BLOCK_SIZE: tl.constexpr,
):
row = tl.program_id(0)
block = tl.program_id(1)
offsets = block * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
values = tl.load(
logits + row * stride_row + offsets * stride_col,
mask=offsets < vocab_size,
other=-float("inf"),
)
values = values.to(tl.float32)

block_max = tl.max(values, axis=0)
block_sum = tl.sum(tl.exp(values - block_max), axis=0)
block_argmax = tl.argmax(values, axis=0) + block * BLOCK_SIZE
output_offset = row * num_blocks + block
tl.store(partial_max + output_offset, block_max)
tl.store(partial_sum + output_offset, block_sum)
tl.store(partial_argmax + output_offset, block_argmax)


@triton.jit
def _greedy_sample_stage2_stats_kernel(
partial_max,
partial_sum,
partial_argmax,
output_stats,
output_argmax,
num_blocks: tl.constexpr,
batch_size: tl.constexpr,
BLOCK_SIZE: tl.constexpr,
):
row = tl.program_id(0)
offsets = tl.arange(0, BLOCK_SIZE)
mask = offsets < num_blocks
input_offset = row * num_blocks + offsets
block_max = tl.load(partial_max + input_offset, mask=mask, other=-float("inf"))
block_sum = tl.load(partial_sum + input_offset, mask=mask, other=0.0)
block_argmax = tl.load(partial_argmax + input_offset, mask=mask, other=0x7FFFFFFF)

global_max = tl.max(block_max, axis=0)
global_sum = tl.sum(block_sum * tl.exp(block_max - global_max), axis=0)
candidate_ids = tl.where(block_max == global_max, block_argmax, 0x7FFFFFFF)
global_argmax = tl.min(candidate_ids, axis=0)
tl.store(output_stats + row, global_max)
tl.store(output_stats + batch_size + row, global_max + tl.log(global_sum))
tl.store(output_argmax + row, global_argmax)


def _launch_stage1(logits: torch.Tensor, scratch: torch.Tensor, block_size: int, num_blocks: int) -> None:
batch_size, vocab_size = logits.shape
_greedy_sample_stage1_kernel[(batch_size, num_blocks)](
logits,
scratch[0],
scratch[1],
scratch[2],
logits.stride(0),
logits.stride(1),
vocab_size=vocab_size,
num_blocks=num_blocks,
BLOCK_SIZE=block_size,
num_warps=8,
)


@torch.no_grad()
def greedy_sample_local_stats(logits: torch.Tensor, alloc_func=torch.empty) -> torch.Tensor:
"""Return local max, logsumexp and argmax rows for distributed greedy sampling."""

assert logits.ndim == 2 and logits.is_cuda and logits.is_contiguous()
batch_size, vocab_size = logits.shape
block_size = 4096
num_blocks = triton.cdiv(vocab_size, block_size)
scratch = alloc_func((3, batch_size, num_blocks), dtype=torch.float32, device=logits.device)
# The third FP32 row carries INT32 argmax bits. Keeping one fixed-size
# payload gives the distributed reducer a single collective without losing
# token-id precision through a numeric int-to-float conversion.
output_stats = alloc_func((3, batch_size), dtype=torch.float32, device=logits.device)

_launch_stage1(logits, scratch, block_size, num_blocks)
_greedy_sample_stage2_stats_kernel[(batch_size,)](
scratch[0],
scratch[1],
scratch[2],
output_stats,
output_stats[2].view(torch.int32),
num_blocks=num_blocks,
batch_size=batch_size,
BLOCK_SIZE=triton.next_power_of_2(num_blocks),
num_warps=4,
)
return output_stats
Original file line number Diff line number Diff line change
@@ -0,0 +1,118 @@
"""Greedy sampling directly from tensor-parallel vocabulary shards."""

import torch
import triton
import triton.language as tl

from lightllm.common.basemodel.triton_kernel.post_process.greedy_sample import (
greedy_sample_local_stats,
)
from lightllm.common.basemodel.triton_kernel.transpose_convert import (
transpose_convert_2d,
)
from lightllm.distributed.communication_op import all_gather_into_tensor


@triton.jit
def _combine_vocab_parallel_stats_kernel(
gathered_stats,
gathered_argmax,
output_logits,
output_token_ids,
output_logsumexp,
token_num,
vocab_size: tl.constexpr,
tp_world_size: tl.constexpr,
BLOCK_SIZE: tl.constexpr,
):
token_offsets = tl.program_id(0) * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
token_mask = token_offsets < token_num
rank_stride = 3 * token_num

global_max = tl.full((BLOCK_SIZE,), -float("inf"), tl.float32)
global_id = tl.full((BLOCK_SIZE,), 0x7FFFFFFF, tl.int32)
for rank in tl.static_range(tp_world_size):
rank_base = rank * rank_stride
local_max = tl.load(
gathered_stats + rank_base + token_offsets,
mask=token_mask,
other=-float("inf"),
)
local_id = tl.load(
gathered_argmax + rank_base + 2 * token_num + token_offsets,
mask=token_mask,
other=0x7FFFFFFF,
)
local_id += (rank * vocab_size) // tp_world_size
wins = (local_max > global_max) | ((local_max == global_max) & (local_id < global_id))
global_max = tl.where(wins, local_max, global_max)
global_id = tl.where(wins, local_id, global_id)

global_sum = tl.zeros((BLOCK_SIZE,), tl.float32)
for rank in tl.static_range(tp_world_size):
rank_base = rank * rank_stride
local_lse = tl.load(
gathered_stats + rank_base + token_num + token_offsets,
mask=token_mask,
other=-float("inf"),
)
global_sum += tl.exp(local_lse - global_max)

tl.store(output_logits + token_offsets, global_max, mask=token_mask)
tl.store(output_token_ids + token_offsets, global_id, mask=token_mask)
tl.store(output_logsumexp + token_offsets, global_max + tl.log(global_sum), mask=token_mask)


@torch.no_grad()
def vocab_parallel_greedy(
local_logits: torch.Tensor,
*,
vocab_size: int,
tp_world_size: int,
group,
alloc_func,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""Return exact sparse logits, global token ids and full-vocab logsumexp."""

assert local_logits.ndim == 2 and local_logits.is_cuda and local_logits.is_contiguous()
local_vocab_size, token_num = local_logits.shape
assert local_vocab_size in {
vocab_size // tp_world_size,
(vocab_size + tp_world_size - 1) // tp_world_size,
}

transposed_logits = alloc_func(
(token_num, local_vocab_size),
dtype=local_logits.dtype,
device=local_logits.device,
)
transpose_convert_2d(local_logits, transposed_logits)
local_stats = greedy_sample_local_stats(transposed_logits, alloc_func=alloc_func)

if tp_world_size == 1:
gathered_stats = local_stats.view(1, 3, token_num)
else:
gathered_stats = alloc_func((tp_world_size, 3, token_num), dtype=torch.float32, device=local_logits.device)
all_gather_into_tensor(
output_=gathered_stats,
input_=local_stats,
group=group,
async_op=False,
)

output_logits = alloc_func((token_num, 1), dtype=torch.float32, device=local_logits.device)
output_token_ids = alloc_func((token_num, 1), dtype=torch.int64, device=local_logits.device)
output_logsumexp = alloc_func((token_num,), dtype=torch.float32, device=local_logits.device)
_combine_vocab_parallel_stats_kernel[(triton.cdiv(token_num, 256),)](
gathered_stats,
gathered_stats.view(torch.int32),
output_logits,
output_token_ids,
output_logsumexp,
token_num,
vocab_size=vocab_size,
tp_world_size=tp_world_size,
BLOCK_SIZE=256,
num_warps=4,
)
return output_logits, output_token_ids, output_logsumexp
Loading
Loading