作者:吴业亮
博客:wuyeliang.blog.csdn.net

本文将详细实现基于LangGraph构建医疗辅助诊断工作流,结合本地VLLM部署医疗大模型(适配A40 48G显卡),全程基于Ubuntu22.04+Conda环境。

一、系统架构

核心架构分为三层:

  1. 模型层:通过VLLM本地部署医疗大模型(如Qwen2-7B-Medical),提供OpenAI兼容API;
  2. 工作流层:LangGraph构建多节点诊断流程(症状收集→初步诊断→检查建议→治疗建议);
  3. 交互层:命令行/WEB界面(Gradio)实现用户交互。

二、环境准备(Ubuntu22.04)

1. 安装Miniconda

# 下载并安装Miniconda
wget https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh
bash Miniconda3-latest-Linux-x86_64.sh
# 生效环境变量
source ~/.bashrc

2. 创建并激活Conda环境

conda create -n med-diagnosis python=3.10 -y
conda activate med-diagnosis

3. 安装显卡驱动与CUDA(适配A40)

A40需NVIDIA驱动≥525,CUDA≥11.8(推荐12.1):

# 安装NVIDIA驱动535(兼容CUDA12.1)
sudo apt update && sudo apt install nvidia-driver-535 -y
sudo reboot  # 重启生效

# 安装CUDA 12.1
wget https://developer.download.nvidia.com/compute/cuda/12.1.0/local_installers/cuda_12.1.0_530.30.02_linux.run
sudo sh cuda_12.1.0_530.30.02_linux.run --override  # 安装时取消Driver(已装),仅选CUDA Toolkit

# 配置CUDA环境变量
echo "export PATH=/usr/local/cuda-12.1/bin:$PATH" >> ~/.bashrc
echo "export LD_LIBRARY_PATH=/usr/local/cuda-12.1/lib64:$LD_LIBRARY_PATH" >> ~/.bashrc
source ~/.bashrc

# 验证
nvcc -V  # 显示CUDA 12.1
nvidia-smi  # 看到A40显卡信息

4. 安装核心依赖

# PyTorch(适配CUDA12.1)
pip3 install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu121

# VLLM(本地大模型推理)
pip install vllm

# LangGraph+LangChain(工作流+API调用)
pip install langgraph langchain langchain-openai pydantic python-dotenv

# 模型下载工具(可选)
pip install modelscope

三、下载医疗大模型

选择轻量且适配医疗场景的Qwen2-7B-Medical(A40 48G无压力运行):

# download_model.py
from modelscope.hub.snapshot_download import snapshot_download

# 下载模型到~/models/Qwen2-7B-Medical
model_dir = snapshot_download(
    model_id="qwen/Qwen2-7B-Medical",
    cache_dir="~/models",
    revision="master"
)
print(f"模型下载完成:{model_dir}")

运行下载:

python download_model.py

四、启动VLLM OpenAI兼容服务器

创建启动脚本start_vllm.sh(适配A40 48G显存):

#!/bin/bash
conda activate med-diagnosis

# 启动VLLM服务器(OpenAI兼容)
vllm serve ~/models/qwen/Qwen2-7B-Medical \
  --host 0.0.0.0 \
  --port 8000 \
  --tensor-parallel-size 1 \  # A40单卡,设为1
  --gpu-memory-utilization 0.9 \  # 占用90%显存(≈43G)
  --max-num-batched-tokens 4096 \
  --temperature 0.1 \  # 医疗场景低温度,结果更稳定
  --max-tokens 2048

赋予权限并启动:

chmod +x start_vllm.sh
./start_vllm.sh  # 保持终端运行,另开终端执行后续步骤

五、LangGraph构建诊断工作流

创建核心脚本med_diagnosis_graph.py

from langchain_openai import ChatOpenAI
from langgraph.graph import StateGraph, END
from pydantic import BaseModel, Field
from typing import Optional, Literal

# 1. 定义诊断状态(存储全流程信息)
class DiagnosisState(BaseModel):
    user_symptoms: str = Field(default="", description="用户症状描述")
    preliminary_diagnosis: Optional[str] = Field(default=None, description="初步诊断")
    examination_suggestions: Optional[str] = Field(default=None, description="检查建议")
    treatment_suggestions: Optional[str] = Field(default=None, description="治疗建议")
    need_further_inquiry: bool = Field(default=True, description="是否需要进一步问诊")

# 2. 连接本地VLLM的OpenAI API
llm = ChatOpenAI(
    base_url="http://localhost:8000/v1",  # VLLM兼容端点
    api_key="dummy_key",  # VLLM无需真实API Key
    model_name="Qwen2-7B-Medical",
    temperature=0.1,
    max_tokens=2048
)

# 3. 定义节点函数
## 节点1:症状收集
def collect_symptoms(state: DiagnosisState) -> DiagnosisState:
    if state.user_symptoms:
        print("\n=== 症状确认 ===")
        print(f"当前症状:{state.user_symptoms}")
        further_info = input("是否补充症状?(输入内容或'无需补充'):")
        if further_info != "无需补充":
            state.user_symptoms += f" {further_info}"
        state.need_further_inquiry = False
    else:
        print("\n=== 医疗辅助诊断系统 ===")
        print("⚠️  免责声明:本系统仅为辅助,不能替代专业医生!")
        state.user_symptoms = input("请详细描述症状(如:发烧38.5℃,咳嗽2天):")
    return state

## 节点2:初步诊断
def preliminary_diagnosis(state: DiagnosisState) -> DiagnosisState:
    print("\n=== 初步诊断 ===")
    prompt = f"""
    你是医疗辅助诊断专家,基于以下症状分析:
    症状:{state.user_symptoms}
    要求:
    1. 列出Top3可能的疾病;
    2. 分析症状与疾病的关联;
    3. 判断是否需要更多症状信息(是/否);
    输出清晰易懂,避免专业术语过载。
    """
    response = llm.invoke(prompt)
    state.preliminary_diagnosis = response.content
    # 简单判断是否需要进一步问诊
    state.need_further_inquiry = "需要进一步" in response.content or "信息不足" in response.content
    print(f"结果:\n{state.preliminary_diagnosis}")
    return state

## 节点3:检查建议
def suggest_examinations(state: DiagnosisState) -> DiagnosisState:
    print("\n=== 检查建议 ===")
    prompt = f"""
    基于初步诊断给出针对性检查建议:
    初步诊断:{state.preliminary_diagnosis}
    要求:
    1. 列出检查项目及目的;
    2. 标注检查优先级;
    输出清晰易懂。
    """
    response = llm.invoke(prompt)
    state.examination_suggestions = response.content
    print(f"结果:\n{state.examination_suggestions}")
    return state

## 节点4:治疗建议
def suggest_treatment(state: DiagnosisState) -> DiagnosisState:
    print("\n=== 治疗建议 ===")
    prompt = f"""
    基于诊断和检查建议给出辅助治疗方案:
    初步诊断:{state.preliminary_diagnosis}
    检查建议:{state.examination_suggestions}
    要求:
    1. 包含药物、生活方式建议;
    2. 强制标注“仅为辅助,需遵医嘱”;
    3. 提醒症状加重及时就医。
    """
    response = llm.invoke(prompt)
    state.treatment_suggestions = response.content
    print(f"结果:\n{state.treatment_suggestions}")
    return state

# 4. 条件判断:是否继续问诊
def should_continue(state: DiagnosisState) -> Literal["collect_symptoms", "suggest_examinations"]:
    return "collect_symptoms" if state.need_further_inquiry else "suggest_examinations"

# 5. 构建LangGraph工作流
def build_graph():
    graph = StateGraph(DiagnosisState)
    # 添加节点
    graph.add_node("collect_symptoms", collect_symptoms)
    graph.add_node("preliminary_diagnosis", preliminary_diagnosis)
    graph.add_node("suggest_examinations", suggest_examinations)
    graph.add_node("suggest_treatment", suggest_treatment)
    # 起始节点
    graph.set_entry_point("collect_symptoms")
    # 边定义
    graph.add_edge("collect_symptoms", "preliminary_diagnosis")
    # 条件边:初步诊断后判断是否继续问诊
    graph.add_conditional_edges(
        "preliminary_diagnosis",
        should_continue,
        {"collect_symptoms": "collect_symptoms", "suggest_examinations": "suggest_examinations"}
    )
    # 检查建议→治疗建议→结束
    graph.add_edge("suggest_examinations", "suggest_treatment")
    graph.add_edge("suggest_treatment", END)
    return graph.compile()

# 6. 运行诊断流程
if __name__ == "__main__":
    diagnosis_graph = build_graph()
    initial_state = DiagnosisState()
    final_state = diagnosis_graph.invoke(initial_state)
    
    # 汇总结果
    print("\n=== 诊断汇总 ===")
    print(f"症状:{final_state.user_symptoms}")
    print(f"初步诊断:{final_state.preliminary_diagnosis}")
    print(f"检查建议:{final_state.examination_suggestions}")
    print(f"治疗建议:{final_state.treatment_suggestions}")
    print("\n⚠️  再次提醒:请及时咨询专业医生!")

六、运行系统

  1. 确保VLLM服务器已启动(./start_vllm.sh);
  2. 运行诊断脚本:
python med_diagnosis_graph.py

示例交互流程:

=== 医疗辅助诊断系统 ===
⚠️  免责声明:本系统仅为辅助,不能替代专业医生!
请详细描述症状(如:发烧38.5℃,咳嗽2天):发烧39℃,头痛,肌肉酸痛,持续1天

=== 初步诊断 ===
结果:
1. 流行性感冒(流感):高烧(39℃)、头痛、肌肉酸痛是流感典型症状,且起病急(1天);
2. 普通感冒:但普通感冒高烧较少见,多为低热;
3. 病毒性脑炎(低概率):头痛伴高烧需警惕,但无呕吐、意识障碍等,暂不优先考虑;
无需进一步获取症状信息。

=== 检查建议 ===
结果:
1. 血常规(优先级:高):判断是否为病毒/细菌感染;
2. 流感病毒核酸检测(优先级:高):明确是否为流感;
3. 体温监测(优先级:中):每4小时测一次,记录变化;

=== 治疗建议 ===
结果:
1. 药物建议:可服用对乙酰氨基酚(泰诺林)退烧止痛,成人每次500mg,每6小时一次,避免空腹;
2. 生活方式:多喝水,保证休息,避免劳累,室内通风;
3. 注意事项:服药3天体温未降需就医;
仅为辅助,需遵医嘱。

=== 诊断汇总 ===
症状:发烧39℃,头痛,肌肉酸痛,持续1天
初步诊断:...
检查建议:...
治疗建议:...

⚠️  再次提醒:请及时咨询专业医生!

七、优化与扩展

1. 模型升级(适配A40 48G)

如需运行更大模型(如Qwen2-72B-Medical),启用4bit量化:

# 修改start_vllm.sh,添加量化参数
vllm serve ~/models/qwen/Qwen2-72B-Medical \
  --host 0.0.0.0 --port 8000 \
  --tensor-parallel-size 1 \
  --gpu-memory-utilization 0.95 \
  --quantization awq \  # 4bit量化
  --max-num-batched-tokens 8192

2. WEB界面(Gradio)

安装Gradio并创建web版:

pip install gradio

创建web_diagnosis.py

import gradio as gr
from med_diagnosis_graph import build_graph, DiagnosisState

diagnosis_graph = build_graph()

def run_web_diagnosis(symptoms):
    initial_state = DiagnosisState(user_symptoms=symptoms)
    final_state = diagnosis_graph.invoke(initial_state)
    return f"""
    ### 初步诊断
    {final_state.preliminary_diagnosis}

    ### 检查建议
    {final_state.examination_suggestions}

    ### 治疗建议
    {final_state.treatment_suggestions}

    ⚠️ 免责声明:本结果仅为辅助参考,请及时咨询专业医生!
    """

with gr.Blocks(title="医疗辅助诊断系统") as demo:
    gr.Markdown("# 医疗辅助诊断系统(仅作辅助)")
    gr.Markdown("⚠️ 本系统不能替代专业医生的诊断!")
    symptoms = gr.Textbox(label="症状描述", lines=5)
    submit = gr.Button("开始诊断")
    output = gr.Markdown(label="诊断结果")
    submit.click(run_web_diagnosis, inputs=symptoms, outputs=output)

if __name__ == "__main__":
    demo.launch(server_name="0.0.0.0", server_port=7860)

运行WEB版:

python web_diagnosis.py

访问http://服务器IP:7860即可使用。

3. 显存监控

实时监控A40显存使用:

watch -n 1 nvidia-smi

八、注意事项

  1. 医疗合规:本系统仅为技术演示,严禁用于临床诊断,需添加明确免责声明;
  2. 模型选择:优先使用医疗领域微调的开源模型,避免通用模型的误诊风险;
  3. 显存管理:A40 48G可运行7B/14B(无量化)、72B(4bit量化)模型;
  4. 性能调优:VLLM的--gpu-memory-utilization可根据显存使用调整(0.8~0.95)。

更多推荐