v01_0_ragged_batch
本版本引入什么
本版本把不同提示词长度的请求词元展平到一个缓冲区,并使用 req_indptr 标记请求边界,从而完成批处理。
模型接收紧凑的不规则布局,不再把每个请求填充到相同长度。
为什么引入
真实推理很少每次只收到一个形状完美的请求。批处理可以提高利用率,但把短提示词填充到最长提示词会浪费计算。不规则批处理展示了怎样在保持各请求因果注意力独立的同时处理变长请求。
核心原则
req_indptr 是前缀和索引。请求 i 拥有以下词元切片:
flat_input_ids[req_indptr[i] : req_indptr[i + 1]]
GQA 遍历这些切片,在每个请求内部执行因果注意力,并按原有展平顺序写回结果。
建议对比的文件
forward_params.py:新增的请求边界元数据。request.py:请求容器。llm.py:批输入构造。layer/gqa.py:在展平缓冲区上逐请求计算注意力。
保留的权衡
代码有意保留逐请求 Python 循环。它不是最快方案,但能在后续引入缓存和调度前,把不规则布局表达清楚。