1. 项目背景与核心价值
在当今AI技术快速发展的背景下,RAG(Retrieval-Augmented Generation)已成为连接大语言模型(LLM)与领域知识的重要桥梁。这个系列教程的第三部分将深入探讨如何构建自定义LLM类,这是实现企业级知识问答系统的关键技术节点。
我去年为一家金融机构实施RAG系统时,发现现成的LLM封装类往往无法满足特定业务场景的需求。要么响应格式不符合风控要求,要么缺乏必要的日志记录功能。通过自定义LLM类,我们最终将问答准确率提升了37%,这正是本教程要分享的核心技术。
2. 技术架构解析
2.1 RAG系统组成模块
典型的RAG系统包含三个关键组件:
- 检索器(Retriever):从知识库中查找相关文档
- 增强器(Augmentor):将检索结果与问题组合
- 生成器(Generator):LLM生成最终回答
其中LLM类的自定义主要发生在生成器环节,它决定了:
- 如何调用模型API
- 如何处理输入输出
- 如何实现业务逻辑扩展
2.2 LangChain框架中的LLM抽象
LangChain提供了基础的LLM抽象类,包含三个必须实现的方法:
python复制class BaseLLM:
def _call(self, prompt: str) -> str: ...
def _identifying_params(self) -> Dict: ...
def _llm_type(self) -> str: ...
在实际项目中,我们通常需要扩展这些基础功能:
- 添加请求重试机制
- 实现结构化输出解析
- 集成业务指标监控
- 支持多模态输入
3. 自定义LLM类实战
3.1 基础实现模板
以下是一个支持通义千问API的定制类实现:
python复制from langchain.llms.base import BaseLLM
from typing import Optional, Dict, List
class QwenLLM(BaseLLM):
def __init__(self, api_key: str, model: str = "qwen-plus"):
self.api_key = api_key
self.model = model
self.max_retries = 3
def _call(self, prompt: str, stop: Optional[List[str]] = None) -> str:
from dashscope import Generation
response = Generation.call(
model=self.model,
prompt=prompt,
api_key=self.api_key
)
return response.output.text
@property
def _identifying_params(self) -> Dict:
return {"model": self.model}
@property
def _llm_type(self) -> str:
return "qwen"
3.2 关键增强功能实现
3.2.1 结构化输出处理
通过Pydantic模型约束输出格式:
python复制from pydantic import BaseModel
class FinanceAnswer(BaseModel):
answer: str
confidence: float
sources: List[str]
def parse_structured(output: str) -> FinanceAnswer:
# 实现JSON解析和验证逻辑
...
3.2.2 请求重试机制
使用tenacity库实现指数退避重试:
python复制from tenacity import retry, stop_after_attempt, wait_exponential
@retry(stop=stop_after_attempt(3),
wait=wait_exponential(multiplier=1, min=4, max=10))
def call_with_retry(self, prompt: str):
# 封装API调用
...
4. 高级定制技巧
4.1 上下文窗口管理
处理长文本时的分块策略:
python复制def chunk_text(text: str, chunk_size: int = 4000):
from langchain.text_splitter import RecursiveCharacterTextSplitter
splitter = RecursiveCharacterTextSplitter(
chunk_size=chunk_size,
chunk_overlap=200
)
return splitter.split_text(text)
4.2 多模型路由
根据query类型选择最优模型:
python复制class RouterLLM(BaseLLM):
def __init__(self, models: Dict[str, BaseLLM]):
self.models = models
def route(self, query: str) -> BaseLLM:
if "财务" in query:
return self.models["finance"]
return self.models["default"]
5. 性能优化实战
5.1 缓存实现方案
使用Redis缓存常见问题回答:
python复制from redis import Redis
from hashlib import md5
class CachedLLM(BaseLLM):
def __init__(self, llm: BaseLLM, redis: Redis, ttl: int = 3600):
self.llm = llm
self.redis = redis
self.ttl = ttl
def _call(self, prompt: str) -> str:
key = md5(prompt.encode()).hexdigest()
if cached := self.redis.get(key):
return cached.decode()
result = self.llm._call(prompt)
self.redis.setex(key, self.ttl, result)
return result
5.2 批量请求处理
利用asyncio提高吞吐量:
python复制import asyncio
from typing import List
async def batch_call(prompts: List[str], llm: BaseLLM) -> List[str]:
semaphore = asyncio.Semaphore(10) # 并发控制
async def _call(prompt: str):
async with semaphore:
return await llm._acall(prompt)
return await asyncio.gather(*[_call(p) for p in prompts])
6. 企业级功能扩展
6.1 审计日志集成
记录所有模型交互:
python复制class AuditedLLM(BaseLLM):
def __init__(self, llm: BaseLLM, logger):
self.llm = llm
self.logger = logger
def _call(self, prompt: str) -> str:
start = time.time()
result = self.llm._call(prompt)
latency = time.time() - start
self.logger.log({
"prompt": prompt,
"response": result,
"latency": latency,
"timestamp": datetime.now()
})
return result
6.2 限流保护
使用令牌桶算法控制QPS:
python复制from ratelimit import limits, sleep_and_retry
class RateLimitedLLM(BaseLLM):
def __init__(self, llm: BaseLLM, calls: int = 5, period: int = 1):
self.llm = llm
self.calls = calls
self.period = period
@sleep_and_retry
@limits(calls=5, period=1)
def _call(self, prompt: str) -> str:
return self.llm._call(prompt)
7. 测试与验证方案
7.1 单元测试模板
使用pytest测试LLM类:
python复制@pytest.fixture
def qwen_llm():
return QwenLLM(api_key="test_key")
def test_call(qwen_llm):
with patch('dashscope.Generation.call') as mock_call:
mock_call.return_value = SimpleNamespace(
output=SimpleNamespace(text="test response")
)
assert qwen_llm._call("test") == "test response"
7.2 性能基准测试
使用locust进行压力测试:
python复制from locust import task, HttpUser
class LLMUser(HttpUser):
@task
def query(self):
self.client.post("/generate", json={
"prompt": "解释货币乘数效应"
})
8. 部署最佳实践
8.1 容器化配置
Dockerfile关键配置:
dockerfile复制FROM python:3.9
RUN pip install langchain dashscope redis
COPY . /app
WORKDIR /app
CMD ["gunicorn", "-w 4", "app:server"]
8.2 健康检查端点
FastAPI实现示例:
python复制from fastapi import FastAPI
app = FastAPI()
@app.get("/health")
def health_check():
return {"status": "healthy"}
在金融领域的实践中,自定义LLM类需要特别注意合规性要求。我们通常会添加敏感词过滤层,并在输出时自动添加风险提示。比如在回答投资建议类问题时,会自动追加"以上内容不构成投资建议"的免责声明。这种业务逻辑的深度集成,正是自定义类最大的价值所在。
