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)

性能优化意义

这样做的主要好处是:

  1. 减少计算量:避免对不需要的token位置计算logits

  2. 节省内存sample_hidden_states 比 hidden_states 小很多

  3. 提高效率:只处理真正需要采样的位置

实际例子

假设批次中有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,避免不必要的计算。

Logo

腾讯云面向开发者汇聚海量精品云计算使用和开发经验,营造开放的云计算技术生态圈。

更多推荐