LoRA 模型运行时热插拔(在线无缝切换)完整可运行代码

一、代码说明

  • 核心能力:大模型主干全程冻结、不重载、不重启服务,推理过程中毫秒级切换不同LoRA适配器

  • 无任何第三方重型依赖:纯PyTorch手写LoRA层,不用transformers、vLLM,剥离框架黑盒,看懂底层原生切换逻辑

  • 线上真实流程复刻:常驻主干模型 + 预加载多套LoRA插件 + 接口动态切换激活LoRA + 并发推理无中断

  • 贴合你之前的疑问:直观展示运行中切换只改指针、不动主干权重、无重新加载耗时

二、完整代码(直接复制运行)

import torch
import torch.nn as nn
from fastapi import FastAPI, UploadFile
import uvicorn
import os
import pickle

# ===================== 1. 原生手写LoRA层(底层核心,和LLM官方实现一致) =====================
class LoRALinear(nn.Module):
    def __init__(self, in_dim: int, out_dim: int, lora_rank: int = 8, alpha: float = 8):
        super().__init__()
        # 冻结主干大权重:全程不更新、不重新加载、永远常驻显存
        self.base_weight = nn.Linear(in_dim, out_dim, bias=False)
        for param in self.base_weight.parameters():
            param.requires_grad = False

        # LoRA双低秩矩阵:仅这部分是可切换的小权重
        self.lora_A = nn.Linear(in_dim, lora_rank, bias=False)
        self.lora_B = nn.Linear(lora_rank, out_dim, bias=False)
        self.scale = alpha / lora_rank

    def forward(self, x):
        # 主干固定输出 + LoRA增量输出
        base_out = self.base_weight(x)
        lora_out = self.lora_B(self.lora_A(x)) * self.scale
        return base_out + lora_out

    # 关键接口:运行时替换LoRA权重,主干不动
    def load_lora_weight(self, new_lora_A, new_lora_B):
        # 仅覆盖两个小矩阵,毫秒级完成,不碰主干base_weight
        self.lora_A.weight.data.copy_(new_lora_A)
        self.lora_B.weight.data.copy_(new_lora_B)

# ===================== 2. 模拟大模型主干网络(Transformer注意力线性层) =====================
class MockLLM(nn.Module):
    def __init__(self):
        super().__init__()
        # 模拟LLM Q/K/V注意力层,接入LoRA
        self.q_proj = LoRALinear(in_dim=512, out_dim=512, lora_rank=8)
        self.k_proj = LoRALinear(in_dim=512, out_dim=512, lora_rank=8)

    def forward(self, x):
        q = self.q_proj(x)
        k = self.k_proj(x)
        return q + k

    # 全局切换LoRA:统一更换所有层的LoRA插件
    def switch_all_lora(self, lora_dict):
        self.q_proj.load_lora_weight(lora_dict["A"], lora_dict["B"])
        self.k_proj.load_lora_weight(lora_dict["A"], lora_dict["B"])
        return True

# ===================== 3. 预生成3套不同LoRA权重(模拟线上训练好的不同场景LoRA) =====================
def build_lora_bank():
    # LoRA1:通用对话场景
    lora_chat = {
        "A": torch.randn(8, 512) * 0.01,
        "B": torch.randn(512, 8) * 0.01
    }
    # LoRA2:代码生成场景
    lora_code = {
        "A": torch.randn(8, 512) * 0.05,
        "B": torch.randn(512, 8) * 0.05
    }
    # LoRA3:文案创作场景
    lora_writer = {
        "A": torch.randn(8, 512) * 0.1,
        "B": torch.randn(512, 8) * 0.1
    }
    return {
        "chat": lora_chat,
        "code": lora_code,
        "writer": lora_writer
    }

# ===================== 4. 启动常驻模型服务(核心:主干模型只加载一次,永久常驻) =====================
app = FastAPI(title="LoRA在线热切换&动态新增更新服务")

# 全局常量:LoRA权重持久化存储目录
LORA_SAVE_DIR = "./lora_storage"
os.makedirs(LORA_SAVE_DIR, exist_ok=True)

# 1. 全局初始化:主干模型仅加载一次,全程不销毁、不重载
llm_model = MockLLM()
# 2. 提前加载全部LoRA插件到内存缓存池
lora_bank = build_lora_bank()
# 3. 全局指针:记录当前正在使用的LoRA(切换本质就是修改这个指针)
current_lora_key = "chat"
llm_model.switch_all_lora(lora_bank[current_lora_key])

print("✅ LLM主干模型加载完成,永久常驻显存")
print("✅ 初始LoRA插件预加载完成:chat / code / writer")
print("✅ 当前默认激活LoRA:通用对话chat")
print("✅ 支持不停机新增LoRA、覆盖更新已有LoRA、磁盘持久化保存")

# ===================== 5. 新增:三大核心接口(适配后期新增/更新LoRA需求) =====================
# 接口1:原有推理接口(不变)
@app.post("/infer")
def model_infer():
    # 模拟模型输入
    x = torch.randn(1, 512)
    out = llm_model(x)
    return {
        "code": 200,
        "current_lora": current_lora_key,
        "model_output_shape": list(out.shape),
        "msg": "推理成功,服务无中断"
    }

# 接口2:原有LoRA热切换接口(不变)
@app.post("/switch_lora")
def switch_lora(lora_type: str):
    global current_lora_key
    if lora_type not in lora_bank.keys():
        return {"code": 400, "msg": f"仅支持切换:{list(lora_bank.keys())}"}
    
    # 一行代码完成热切换:仅替换LoRA小权重,主干模型完全不动
    llm_model.switch_all_lora(lora_bank[lora_type])
    current_lora_key = lora_type

    return {
        "code": 200,
        "now_active_lora": current_lora_key,
        "switch_cost": "<1ms",
        "msg": "LoRA热切换完成,服务无重启、无中断"
    }

# ---------------------- 【核心新增1】不停机新增全新LoRA权重 ----------------------
@app.post("/add_new_lora")
def add_new_lora(lora_name: str):
    """线上不停机新增一套全新LoRA,加入内存缓存池,后续可直接切换"""
    global lora_bank
    if lora_name in lora_bank:
        return {"code": 400, "msg": f"LoRA:{lora_name}已存在,请勿重复新增"}
    
    # 模拟新训练完成的LoRA权重(生产环境替换为读取训练产出的lora_A/lora_B)
    new_lora = {
        "A": torch.randn(8, 512) * 0.08,
        "B": torch.randn(512, 8) * 0.08
    }
    # 加入内存缓存池,无需重启服务
    lora_bank[lora_name] = new_lora
    # 持久化保存到磁盘,服务重启后可自动加载
    with open(f"{LORA_SAVE_DIR}/{lora_name}.pkl", "wb") as f:
        pickle.dump(new_lora, f)

    return {
        "code": 200,
        "new_lora_name": lora_name,
        "all_lora_list": list(lora_bank.keys()),
        "msg": "新LoRA动态新增完成,无需重启模型服务,可直接切换使用"
    }

# ---------------------- 【核心新增2】不停机覆盖更新已有LoRA权重 ----------------------
@app.post("/update_exist_lora")
def update_exist_lora(lora_name: str):
    """线上不停机更新已存在的LoRA权重,直接覆盖内存+磁盘,无服务中断"""
    global lora_bank, current_lora_key
    if lora_name not in lora_bank:
        return {"code": 400, "msg": f"LoRA:{lora_name}不存在,无法更新"}
    
    # 模拟迭代训练后的新版LoRA权重(生产环境对接训练脚本输出的最新权重)
    updated_lora = {
        "A": torch.randn(8, 512) * 0.02,
        "B": torch.randn(512, 8) * 0.02
    }
    # 1. 覆盖内存缓存中的旧LoRA
    lora_bank[lora_name] = updated_lora
    # 2. 同步覆盖磁盘持久化文件
    with open(f"{LORA_SAVE_DIR}/{lora_name}.pkl", "wb") as f:
        pickle.dump(updated_lora, f)
    
    # 兼容场景:如果当前正在使用该LoRA,自动实时生效新权重
    if current_lora_key == lora_name:
        llm_model.switch_all_lora(updated_lora)

    return {
        "code": 200,
        "update_lora_name": lora_name,
        "current_active_lora": current_lora_key,
        "msg": "已有LoRA权重热更新完成,内存+磁盘已同步,当前使用该LoRA则立即生效"
    }

# ---------------------- 【核心新增3】从磁盘加载外部新LoRA文件(对接离线训练产出) ----------------------
@app.post("/load_lora_from_file")
async def load_lora_from_file(file: UploadFile, lora_name: str):
    """上传离线训练好的LoRA权重文件,线上直接加载接入,全程不停机"""
    global lora_bank
    content = await file.read()
    new_lora = pickle.loads(content)
    lora_bank[lora_name] = new_lora
    # 持久化落地
    with open(f"{LORA_SAVE_DIR}/{lora_name}.pkl", "wb") as f:
        pickle.dump(new_lora, f)
    return {
        "code": 200,
        "lora_name": lora_name,
        "msg": "外部LoRA权重文件在线加载完成,已加入缓存池"
    }

# ===================== 6. 启动服务 =====================
if __name__ == "__main__":
    # 服务启动后,LLM主干一直运行,支持动态新增、更新、加载LoRA
    uvicorn.run(app, host="0.0.0.0", port=8000)

三、运行&amp;测试步骤(零基础直接操作)

  1. 安装依赖:pip install torch fastapi uvicorn

  2. 运行代码,服务启动后,主干模型永久常驻内存

  3. 调用推理接口,测试正常推理:POST http://127.0.0.1:8000/infer

  4. 在线热切换代码生成LoRA:POST http://127.0.0.1:8000/switch_lora?lora_type=code

  5. 再次调用推理接口,会发现模型输出发生变化,服务全程没有重启、没有卡顿

四、拆解核心切换逻辑(对应线上生产环境)

1. 绝对不动的部分

MockLLM内部所有base_weight主干权重,从服务启动到下线,零修改、零重载、零梯度计算,彻底避免模型崩坏、灾难性遗忘。

2. 真正切换的部分

仅仅覆盖lora_A、lora_B 两个小矩阵,单套LoRA只有几十KB,拷贝耗时低于1毫秒。

3. 和工业框架(vLLM/TGI)对齐点

  • vLLM:维护lora缓存池,全局维护active_lora_id指针

  • 本demo:维护current_lora_key全局变量,逻辑完全一致

  • 生产环境:支持请求级别动态LoRA(不同用户同时用不同LoRA),原理相同