医疗Embedding模型
一、模型介绍
本项目通过微调BGE模型,增强了其对医疗领域的理解,该模型不仅能够应对常见的病情描述, 还能准确处理复杂、模糊的症状,显著提升了科室推荐的准确性。
模型将病情描述转化为向量,并通过存储在 Milvus 向量数据库中的科室信息,检索找到最匹配的科室推荐。
二、使用方法
1、利用训练好的模型对标准科室进行编码,并存入milvus向量数据库中
2、对症状进行编码,并在milvus中检索并返回推荐的科室
import torch
from transformers import AutoModel, AutoTokenizer
import pandas as pd
import numpy as np
# 检查 GPU 是否可用
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
# 加载safetensors格式的模型
model = AutoModel.from_pretrained("./Medical_Embedding_Model", use_safetensors=True).to(device)
# 加载模型的分词器
tokenizer = AutoTokenizer.from_pretrained("./Medical_Embedding_Model")
embedding_dim = 1024
def embed(text):
# 返回一个向量,比如从预训练模型中获得的
encoded_input = tokenizer(text, return_tensors='pt').to(device) # 将输入数据移动到 GPU
with torch.no_grad():
model.half()
model.eval()
output = model(**encoded_input)
embedding = output.pooler_output.cpu().numpy().astype(np.float32).flatten()
return embedding # 请根据实际模型调整
# 连接到Milvus
from pymilvus import connections, FieldSchema, CollectionSchema, DataType, Collection
connections.connect("default", host="localhost", port="19530")
# 定义Collection的Schema
fields = [
FieldSchema(name="id", dtype=DataType.INT64, is_primary=True, auto_id=True),
FieldSchema(name="embedding", dtype=DataType.FLOAT_VECTOR, dim=embedding_dim),
FieldSchema(name="text", dtype=DataType.VARCHAR, max_length=100),
]
schema = CollectionSchema(fields, "department_collection")
# 创建Collection
collection = Collection("department_collection", schema)
# 生成标准科室的embeddings
embeddings = [embed(text) for text in texts]
# 插入数据
entities = [
embeddings,
texts
]
collection.insert(entities)
# 创建索引
index_params = {
"index_type": "IVF_FLAT",
"params": {"nlist": 128},
"metric_type": "L2"
}
collection.create_index("embedding", index_params)
collection.load()
def query(query_text):
query_vector = embed(query_text)
# 搜索
search_params = {
"metric_type": "L2",
"params": {"nprobe": 10},
}
results = collection.search(
[query_vector],
anns_field="embedding",
param=search_params,
limit=3, # 返回10个最相似的结果
expr=None,
output_fields=["text"] # 添加这个字段以便检索时返回text
)
# 打印结果
print("\n\n症状描述:",query_text)
for i,result in enumerate(results[0]):
print(f"推荐挂号科室{i+1}: {result.entity.get('text')}")
query_text = """
胸痛,尤其是在夜间加重,伴有咳嗽、咳痰,有时咳出带血痰,体重明显下降
"""
query(query_text)
三、训练方法
1、使用科室数据微调原始模型得到model_1(使模型理解部分特殊科室之间的联系即可,出现灾难性遗忘没有很大影响)
2、使用model_1和原始模型按照7:3进行混合得到model_2
3、使用特殊数据对model_2微调得到model_3(特殊数据包括部分效果不太好的科室数据,例如:血液透析科等)
4、使用model_3和model_2按照7:3进行混合得到model_4
5、使用医生擅长数据对model_4微调得到model_5(注意去除ICU和内科的数据,否则会因为这些科室中的医生包含大量其他科室医生的擅长从而产生干扰)
6、使用model_5和model_4按照6:4进行混合得到model_6
四、效果展示
病情挂号测试:
输入示例1:
query_text = """
眼睑、手脚或下肢浮肿,尤其是早晨时症状明显
"""
query(query_text)
输出示例1:
症状描述:
眼睑、手脚或下肢浮肿,尤其是早晨时症状明显
推荐挂号科室1: 内科_肾病科
推荐挂号科室2: 内科
推荐挂号科室3: 其他科室
输入示例2:
query_text = """
顽固的高血压,伴有剧烈的头痛、视力模糊,偶尔有剧烈胸痛和呼吸困难,近期还出现了腰背部剧痛,并伴有尿液中带血,体重减轻。
"""
query(query_text)
输出示例2:
症状描述:
顽固的高血压,伴有剧烈的头痛、视力模糊,偶尔有剧烈胸痛和呼吸困难,近期还出现了腰背部剧痛,并伴有尿液中带血,体重减轻。
推荐挂号科室1: 内科
推荐挂号科室2: 冠心病监护病房(CCU)
推荐挂号科室3: 内科_心血管内科
输入示例3:
query_text = """
持续的胸痛、咳嗽,伴有声音嘶哑和吞咽困难,近期有呼吸困难加重,出现脸部浮肿和紫绀,偶尔还会有手指末端发青,近来体重急剧下降
"""
query(query_text)
输出示例3:
症状描述:
持续的胸痛、咳嗽,伴有声音嘶哑和吞咽困难,近期有呼吸困难加重,出现脸部浮肿和紫绀,偶尔还会有手指末端发青,近来体重急剧下降
推荐挂号科室1: 内科_呼吸内科
推荐挂号科室2: 重症医学科(ICU)
推荐挂号科室3: 内科
输入示例4:
query_text = """
胸痛,尤其是在夜间加重,伴有咳嗽、咳痰,有时咳出带血痰,体重明显下降
"""
query(query_text)
输出示例4:
症状描述:
胸痛,尤其是在夜间加重,伴有咳嗽、咳痰,有时咳出带血痰,体重明显下降
推荐挂号科室1: 内科_呼吸内科
推荐挂号科室2: 内科
推荐挂号科室3: 外科_胸外科
五、更多测试数据示例
更多测试数据示例见Disease_description.csv
- Downloads last month
- 8
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support