AI模型缓存:推理不够用了之后
前阵子接了个私活,帮一家创业公司搞AI客服。
等到用户量稍微上来点,问题就炸了——响应时间从3秒飙到了15秒,GPU显存吃满,Redis还时不时报个 OOM,老板在群里@我,那滋味,懂的都懂。
缓存到底在哪存着?
我们要先搞清楚一个概念:AI的缓存不是只有一种。
很多时候我们说的"缓存",指的是"存下用户的输入和模型的输出"。这种叫结果缓存,就像是你把作业抄下来,下次老师再问直接交。
但在模型推理层面,还有一个更底层的东西,叫 KV Cache。这玩意儿是在模型计算过程中产生的,通俗点说,模型在生成第N个字的时候,需要记住前面N-1个字的注意力权重。这些权重(Key和Value矩阵)如果不存下来,每次生成新token都要从头算一遍,那性能简直是灾难。
所以,我们的优化得分两层看:
- 推理层:让模型自己跑得快(KV Cache)。
- 应用层:别让模型瞎跑(结果缓存、语义缓存)。
第一层:KV Cache —— 模型的"短期记忆"
这事儿最直观的体现就是 vLLM 这类推理框架的崛起。
以前我们用 transformers 原生推理,默认其实就开了KV Cache(use_cache=True),但这只是"有",并不代表"好"。
踩坑记录:显存爆炸
我一开始直接上 HuggingFace 的 pipeline,开了缓存,结果对话稍微长一点(比如上下文超过4k tokens),显存直接溢出。
from transformers import AutoModelForCausalLM, AutoTokenizer
model_id = "Qwen/Qwen2.5-7B-Instruct"
tokenizer = AutoTokenizer.from_pretrained(model_id)
model = AutoModelForCausalLM.from_pretrained(
model_id,
device_map="auto",
torch_dtype="auto",
use_cache=True # 开启KV Cache
)
# 然后就是一顿猛操作,直到...
# RuntimeError: CUDA out of memory.
问题在于:KV Cache 是随着生成长度动态增长的。原生的实现管理得比较粗糙,显存碎片严重。
解决方案:上 vLLM
后来咬牙上了 vLLM。这玩意儿的核心就是 PagedAttention,把 KV Cache 像操作系统管理内存一样,切成一页一页的,不仅能解决碎片问题,还能做 Continuous Batching。
# 安装
pip install vllm
# 启动服务,指定 GPU 显存利用率和 KV Cache 块大小
python -m vllm.entrypoints.api_server \
--model Qwen/Qwen2.5-7B-Instruct \
--gpu-memory-utilization 0.9 \
--block-size 16 \
--max-model-len 32768
效果:同样一块 24G 显存的 4090,以前只能并发处理2个长上下文请求,现在能抗住8-10个,吞吐量直接翻倍。
但这也有坑:vLLM 对冷启动不太友好,第一次加载模型巨慢,而且如果你 request 的 max_tokens 设置得太大,它会预先占用大量 KV Cache 空间,导致其他请求进不来。千万别贪心,max_tokens 设成实际需要的2倍顶天了。
第二层:Prompt Cache —— 赖上系统提示词
有时候你会发现,虽然用户问的问题不一样,但 System Prompt(系统提示词)总是那几千字不变的。
比如我的客服场景,System Prompt 里写了一堆公司规章制度、产品手册,这玩意儿每次都要喂给模型算一遍,纯属浪费算力。
现在的模型厂商(比如 Anthropic)或者框架(如 SGLang)都支持了 Prompt Caching。
实战:Anthropic 的 API
import anthropic
client = anthropic.Anthropic()
# 巨大的 System Prompt
system_prompt = """
你是一个专业的客服助手...(此处省略2000字)...
请遵守以下规则...(此处省略1000字)...
"""
response = client.messages.create(
model="claude-3-5-sonnet-20241022",
max_tokens=1024,
system=[
{
"type": "text",
"text": system_prompt,
# 关键在这里:标记这段缓存
"cache_control": {"type": "ephemeral"}
}
],
messages=[{"role": "user", "content": "你好"}]
)
实测数据:在没有 Prompt Cache 的情况下,每次请求都要处理这几千字的 System Prompt,首字延迟(TTFT)大概在 1.2s 左右。开了缓存后,TTFT 直接降到了 0.4s,而且这部分是不收输入 Token 费用的。
坑点:这个缓存是会话级的,如果你用同一个 API Key 但换了不同的 System Prompt,缓存命中率会掉得很难看。另外,千万别在缓存块里塞动态内容(比如当前时间、用户ID),否则缓存直接失效,全是无用功。
第三层:语义缓存 —— “那个意思差不多就行”
这是最"软"的一层缓存。
结果缓存很简单:Key 是用户的问题,Value 是模型回答。Key 必须完全一样才能命中。但用户问问题哪有那么标准?“怎么退货"和"怎么申请退换货”,意思一样,结果缓存就傻眼了。
这时候就需要 语义缓存:把问题转向量,算余弦相似度。
抄作业:Redis + FAISS
我试过纯 Milvus,太重了;纯 Redis 的 FT.SEARCH 向量搜索,性能还行但配置麻烦。最后用了 FAISS 做索引,Redis 存原始数据,简单粗暴。
import faiss
import numpy as np
from sentence_transformers import SentenceTransformer
# 1. 加载编码模型(这个也要占显存,别和推理模型挤同一张卡)
encoder = SentenceTransformer('paraphrase-multilingual-MiniLM-L12-v2')
# 2. 初始化 FAISS 索引
dimension = 384 # 模型维度
index = faiss.IndexFlatIP(dimension) # 内积相似度
# 3. 模拟一个缓存库
cache_store = {} # {idx: {"question": "...", "answer": "..."}}
def get_cache(question):
query_vector = encoder.encode([question])
faiss.normalize_L2(query_vector) # 归一化
# 搜索最相似的 Top 1
scores, indices = index.search(query_vector, 1)
best_score = scores[0][0]
best_idx = indices[0][0]
# 阈值很关键,0.85 是我试出来的,低了容易答非所问,高了命中率上不去
if best_score > 0.85:
return cache_store.get(best_idx)
return None
def save_cache(question, answer):
vector = encoder.encode([question])
faiss.normalize_L2(vector)
idx = len(cache_store)
index.add(vector)
cache_store[idx] = {"question": question, "answer": answer}
踩坑:阈值是个玄学
我最开始设了 0.9,结果命中率感人,只有 10%。后来改成 0.8,命中率上去了,但客服群炸了——模型居然给问"怎么退款"的用户返回了"怎么开发票"的答案。
经验:
- 阈值要分场景:简单问答(如FAQ)阈值可以设低点,0.8 左右;涉及具体业务操作(退款、改密码),必须设到 0.85 甚至 0.9。
- 加个兜底机制:命中的结果不要直接返回,先丢给模型做一个校验:“用户问的是[问题],缓存里有[答案],这个答案合适吗?"。多这一步,准确率能提升一大截,但延迟也加上去了,看你怎么权衡。
第四层:结果缓存 —— 别傻傻全存
这是最基础的一层,用 Redis 或者 Memcached 就行。
import redis
import hashlib
import json
r = redis.Redis(host='localhost', port=6379, db=0)
def get_response_cache(question, model_config):
# 生成唯一的 Key
key_data = f"{question}:{json.dumps(model_config, sort_keys=True)}"
cache_key = hashlib.md5(key_data.encode()).hexdigest()
cached = r.get(cache_key)
if cached:
return json.loads(cached)
return None
def set_response_cache(question, model_config, answer, ttl=3600):
key_data = f"{question}:{json.dumps(model_config, sort_keys=True)}"
cache_key = hashlib.md5(key_data.encode()).hexdigest()
r.setex(cache_key, ttl, json.dumps(answer))
踩坑:TTL(过期时间)怎么设?
我之前直接把 TTL 设成 7天,结果有一次产品更新了话术(把"亲"改成了"您好”),结果 Redis 里全是旧答案,用户又投诉了。
做法:
- 业务相关内容(如价格、政策):TTL 设短点,10分钟到1小时。
- 通用知识(如"怎么使用Python"):TTL 可以长点,24小时甚至更长。
- 手动失效机制:发版的时候,带个脚本
redis-cli FLUSHDB,或者给特定 Key 加上版本号。
终极策略:分层拦截
实战中,我是这么组合的:
- L1 结果缓存:全量匹配,直接返回。(命中率 ~15%)
- L2 语义缓存:向量检索相似度 > 0.85。(命中率 ~30%)
- L3 Prompt Cache:如果命中了 System Prompt 缓存,只补用户问题部分走推理。
- L4 KV Cache (vLLM):在推理过程中复用计算结果。
效果:整体平均响应时间从 5s 压到了 1.2s,P99 延迟从 15s 压到了 3s。成本降了一半,因为直接走缓存的请求根本没调用模型。
最后唠两句
缓存这东西,说白了就是用空间换时间,用一致性换性能。
你不可能既要又要还要。比如你想缓存命中率高,就得接受偶尔回答稍微偏一点;你想响应速度快,就得忍受架构变复杂。
折腾到现在,我最大的感受是:别迷信框架,先看监控。有时候你加了那么多层缓存,结果发现瓶颈压根不在模型推理上,而是在网络IO或者数据库查询上,那才是真·尴尬。
写代码就像过日子,精打细算才能过得长久。但该花的地方还得花,别为了省那几百毫秒的显存,把用户体验给省没了。
AI 路漫漫,缓存作舟楫。祝各位都能在"等模型"的日子里,少一点焦虑,多一点掌控。
版权声明: 本文首发于 指尖魔法屋-AI模型缓存:推理不够用了之后(https://blog.thinkmoon.cn/post/ai-model-caching-optimization/) 转载或引用必须申明原指尖魔法屋来源及源地址!
评论
使用 GitHub 账号登录后即可留言,支持 Markdown。