ITADN

[OCR Phase3] 实现 KV Cache 和 Flash Attention 优化

#2378Openmessere1 创建于 2026-01-06
M
messere1commented
## 任务描述 实现 KV Cache 优化和 Flash Attention 集成,提升 Transformer 模型的推理效率和吞吐量。 **关联主 Issue**: #2348 Phase 3: 模型微调与优化 ## 任务目标 - [ ] 实现 KV Cache 机制,避免重复计算 - [ ] 集成 Flash Attention 2.0(如果硬件支持) - [ ] 实现多批次推理时的 Cache 复用 - [ ] 添加 Cache 管理和清理机制 - [ ] 优化长文本 OCR 场景的内存使用 ## 技术方案 ### 1. KV Cache 实现 - 在模型 `generate()` 时启用 `use_cache=True` - 缓存 past_key_values,避免重复计算 attention - 适用于自回归生成场景 ### 2. Flash Attention 集成 - 检测硬件支持(需要 CUDA 7.5+, Ampere 架构+) - 使用 `attn_implementation="flash_attention_2"` - 降低 attention 计算的显存占用(O(N) vs O(N²)) ### 3. Batch 推理优化 - 实现 batch 内 KV Cache 共享 - padding 策略优化(left padding for generation) - dynamic batching 支持 ### 4. Cache 管理 - 实现 LRU Cache 清理策略 - 添加 Cache 大小限制配置 - 提供 Cache 统计和监控接口 ## 测试要求 ### 功能测试 - [ ] 测试 KV Cache 开启/关闭的功能正确性 - [ ] 测试 Flash Attention 在支持的硬件上正常工作 - [ ] 测试 batch 推理的 Cache 复用 - [ ] 测试 Cache 清理机制 ### 性能测试 - [ ] 对比 KV Cache 开启前后的推理速度 - [ ] 对比 Flash Attention 的显存占用 - [ ] 测试不同 batch size 下的吞吐量 - [ ] 长文本场景(> 2048 tokens)的性能测试 ### 基准测试 - 单图推理延迟(毫秒) - Batch 推理吞吐量(图片/秒) - 显存占用峰值(GB) - 不同输入长度的性能曲线 ## 验收标准 - ✅ KV Cache 启用后,推理速度提升 20-30% - ✅ Flash Attention 显存占用降低 30-40%(长序列) - ✅ Batch=4 时吞吐量提升 2.5-3x - ✅ 长文本(> 2048 tokens)推理不 OOM - ✅ 提供详细的性能对比报告 ## 依赖项 - flash-attn >= 2.0 (可选,需要 CUDA 支持) - transformers >= 4.37.0 (KV Cache 原生支持) - torch >= 2.0 ## 注意事项 - Flash Attention 仅在特定硬件上可用(NVIDIA Ampere+) - NPU 可能不支持 Flash Attention,需要降级策略 - KV Cache 会增加内存占用,需要合理配置上限 ## 优先级 **P1** - 重要优化,显著提升推理性能 ## 预计工时 5-7 天
0 条评论