005、模型集成:把本地大模型塞进Flask应用

昨天深夜调试时遇到个典型问题:同事在Flask路由里直接加载7B参数的模型,每次请求都重新读一遍权重文件。结果第一个请求等了三分半,服务器内存直接飙到32G——典型的“把实验代码当生产代码用”。今天咱们就聊聊怎么把本地大模型妥帖地集成到Flask应用里,避开这些新手陷阱。

模型加载的坑别踩第二次

先看这段问题代码:

@app.route('/generate', methods=['POST'])
def generate():
    # 灾难写法!每次请求都加载模型
    model = torch.load('llama-7b.bin')
    input_text = request.json['text']
    output = model.generate(input_text)
    return jsonify({'result': output})

这种写法在开发阶段可能勉强能跑,但上线就是灾难。大模型权重文件通常几个GB起步,加载到GPU还得转换格式,没个一两分钟下不来。更糟的是,多个请求同时进来会重复加载模型,内存不炸才怪。

正确的模型托管姿势

核心思路就一个:应用启动时加载模型,后续请求共享模型实例。Flask的before_first_request装饰器用在这里正合适:

class ModelManager:
    _instance = None  # 单例模式,防止重复初始化
    
    def __init__(self):
        if ModelManager._instance is not None:
            raise RuntimeError('别重复初始化,用get_instance()')
        
        self.model = None
        self.tokenizer = None
        self.device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
        
    @classmethod
    def get_instance(cls):
        if cls._instance is None:
            cls._instance = cls()
        return cls._instance

# Flask应用工厂里初始化
def create_app():
    app = Flask(__name__)
    model_mgr = ModelManager.get_instance()
    
    @app.before_first_request
    def load_model():
        # 这里会阻塞启动,但只阻塞一次
        print('开始加载模型,去泡杯咖啡吧...')
        model_mgr.model = AutoModelForCausalLM.from_pretrained(
            './models/llama-7b',
            torch_dtype=torch.float16,
            device_map='auto'
        )
        model_mgr.tokenizer = AutoTokenizer.from_pretrained('./models/llama-7b')
        print('模型加载完毕!')

注意那个device_map='auto'参数,它能自动把模型层分布到多GPU上。如果你只有一张卡,改成device_map={'':0}明确指定。

请求处理要加锁

模型推理时有个隐藏问题:大多数模型前向传播不支持并发调用。如果两个请求同时调用model.generate(),轻则输出乱码,重则CUDA报错。得加个线程锁:

from threading import Lock

class ModelManager:
    def __init__(self):
        # ... 其他初始化代码
        self.inference_lock = Lock()
        
    def generate_text(self, prompt, max_length=100):
        with self.inference_lock:  # 确保同时只有一个推理在进行
            inputs = self.tokenizer(prompt, return_tensors='pt').to(self.device)
            with torch.no_grad():  # 省内存关键!
                outputs = self.model.generate(
                    **inputs,
                    max_length=max_length,
                    temperature=0.7,
                    do_sample=True
                )
            return self.tokenizer.decode(outputs[0], skip_special_tokens=True)

那个torch.no_grad()上下文管理器务必加上,不然推理过程中的梯度计算会吃掉额外显存。曾经有次忘记加,16G显存跑7B模型直接OOM。

路由层要做输入校验

直接拿用户输入喂给模型太危险了。有人可能发个几万字的文本,或者带特殊字符的恶意输入:

@app.route('/generate', methods=['POST'])
def generate():
    data = request.get_json()
    
    # 防御性编程
    if not data or 'text' not in data:
        return jsonify({'error': '需要text字段'}), 400
    
    prompt = data['text'].strip()
    if len(prompt) > 2000:
        return jsonify({'error': '输入太长,请限制在2000字符内'}), 400
    if len(prompt) == 0:
        return jsonify({'error': '输入不能为空'}), 400
    
    # 模型推理
    try:
        model_mgr = ModelManager.get_instance()
        result = model_mgr.generate_text(prompt)
        return jsonify({'result': result})
    except RuntimeError as e:
        # CUDA内存不足的常见错误
        if 'CUDA out of memory' in str(e):
            return jsonify({'error': '服务器忙,请稍后重试'}), 503
        raise

输入长度限制不是随便定的。7B模型在4096上下文长度下,处理2000字符的输入大概需要1.5G显存。你得根据自己显卡容量来调整这个阈值。

生产环境必备的优化

如果真要把服务部署出去,还有几个关键点:

批处理支持:很多请求可以攒一起推理,显存利用率能翻倍。但Flask默认不支持,得改造成异步或者用消息队列。简单做法是用concurrent.futures搞个线程池:

from concurrent.futures import ThreadPoolExecutor

executor = ThreadPoolExecutor(max_workers=2)  # 根据GPU数量调整

@app.route('/generate', methods=['POST'])
def generate():
    # ... 校验代码
    future = executor.submit(model_mgr.generate_text, prompt)
    try:
        result = future.result(timeout=30)  # 设置超时
        return jsonify({'result': result})
    except TimeoutError:
        return jsonify({'error': '生成超时'}), 504

显存清理:长时间运行后显存会碎片化,可以定期清理:

import gc
import torch

def cleanup_memory():
    gc.collect()
    if torch.cuda.is_available():
        torch.cuda.empty_cache()
        torch.cuda.ipc_collect()

这个函数可以放在定时任务里,每小时跑一次。

调试小技巧

模型集成阶段最容易遇到“本地能跑,服务器挂掉”的情况。建议在app.run()前加这几行:

if __name__ == '__main__':
    # 先预加载模型,检查配置
    mgr = ModelManager.get_instance()
    mgr.load_model()
    
    # 试运行一次推理
    test_output = mgr.generate_text('Hello')
    print(f'测试输出: {test_output}')
    
    # 确认没问题再启动服务
    app.run(host='0.0.0.0', port=5000, threaded=True)

threaded=True很重要,Flask默认是单线程,同时只能处理一个请求。

个人经验之谈

集成大模型到Web服务,本质是在“开发便利性”和“生产稳定性”之间找平衡。我习惯分三步走:先用最简单的方式跑通流程,然后加上异常处理和资源管理,最后考虑性能优化。别一开始就追求完美,很多优化手段在模型规模确定前都是白费功夫。

有个容易忽略的点:模型文件别放在项目目录里。最好用符号链接到单独的数据盘,这样更新权重时不用重启整个应用。另外,务必给Flask加上访问日志,记录每个请求的输入长度和推理时间,这些数据后期优化时非常有用。

最后说个血泪教训:永远在代码里写死显存分配策略。曾经有次更新驱动后,CUDA的默认分配策略变了,服务直接瘫痪。现在我都显式设置PYTORCH_CUDA_ALLOC_CONF=max_split_size_mb:128,虽然可能损失一点性能,但稳定性提升值得。

模型集成这块,多踩几次坑就熟练了。下次我们聊聊怎么给这个API加上流式输出——总不能等30秒生成完了才返回第一个字吧?

Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐