vllm中生成token的流程
hidden_states 的类型和结构
hidden_states 是模型前向传播的输出,具体来说:
-
类型:
torch.Tensor -
形状:
[total_num_tokens, hidden_size] -
内容:每个输入token对应的最后一层隐藏状态表示
示例:
-
如果有3个请求,token数量分别为 [2, 3, 1],总共6个token
-
hidden_states的形状就是[6, hidden_size](如[6, 4096])
为什么要用 logits_indices 进行索引
核心原因:只需要部分位置的hidden states来计算logits
在序列生成中:
-
我们输入了多个token(包括历史token和新token)
-
但只需要在特定位置生成下一个token
-
不是所有位置的hidden states都需要用于计算logits
具体场景分析
1. 常规解码(非推测解码)
python
# 假设有3个请求,token数量: [2, 3, 1] query_start_loc = [0, 2, 5, 6] # 每个请求的起始位置 logits_indices = [1, 4, 5] # 每个序列的最后一个token位置 hidden_states形状: [6, hidden_size] sample_hidden_states = hidden_states[[1, 4, 5]] # 形状: [3, hidden_size]
为什么取最后一个token?
-
请求1:输入token [t0, t1] → 需要t1的hidden state来生成t2
-
请求2:输入token [t0, t1, t2] → 需要t2的hidden state来生成t3
-
请求3:输入token [t0] → 需要t0的hidden state来生成t1
2. 推测解码(Speculative Decoding)
在推测解码中更复杂,可能需要在多个位置计算logits来验证draft tokens。
完整的数据流
python
# 1. 模型前向传播 hidden_states = model(input_ids, ...) # [total_tokens, hidden_size] # 2. 选择需要计算logits的位置 sample_hidden_states = hidden_states[logits_indices] # [num_requests, hidden_size] # 3. 计算logits(只对选中的位置) logits = self.model.compute_logits(sample_hidden_states) # [num_requests, vocab_size] # 4. 采样 next_tokens = sampler(logits)
性能优化意义
这样做的主要好处是:
-
减少计算量:避免对不需要的token位置计算logits
-
节省内存:
sample_hidden_states比hidden_states小很多 -
提高效率:只处理真正需要采样的位置
实际例子
假设批次中有2个请求:
-
请求1:已有3个token,新增1个token → 总共4个token输入
-
请求2:已有1个token,新增2个token → 总共3个token输入
text
hidden_states形状: [7, hidden_size] # 4 + 3 = 7个token logits_indices = [3, 6] # 请求1的第3个位置,请求2的第2个位置 sample_hidden_states形状: [2, hidden_size] # 大大减小!
这样,系统只需要计算2个位置的logits,而不是7个位置,显著提高了效率。
总结:logits_indices 的作用就是精确定位哪些token的hidden states需要用于生成下一个token,避免不必要的计算。
更多推荐
所有评论(0)