Chuyển đến nội dung chính

レッスン 8: 医療 Q&A と医療チャットボット

医学的な質問に答えました。ヘルスケア向け RAG: PubMed 検索。医療ドメイン向けに LLM を微調整します。ガードレール、医療チャットボットの安全性。

🧠 AI と ML — レッスン 7 レッスン 8: 医療 Q&A と医療チャットボット

医療とヘルスケアにおける AI: 実戦アプリケーション

パート 3: 臨床 NLP とゲノミクス AI

xdev.asia

平均的な医師は知識を更新するために年間 5,000 件の論文を読みます。 AI は適切に構築されていれば、PubMed 上の 3,400 万件の論文を読み、数秒で応答できます。


1. 医療Q&Aシステムの仕組み

User query → Safe query check → Retriever → Context ranking
                                  ↓
              [PubMed / Clinical guidelines / EHR policies]
                                  ↓
                           LLM Generator
                                  ↓
              Response → Safety guardrails → User

LLM だけでなく RAG を使用する理由

  • LLM 幻覚: 存在しない引用を自動的に生成します
  • 医学知識には特定の情報源(PMID、ガイドライン)を引用する必要があります
  • モデルを再トレーニングせずに知識を更新する必要がある

2. 検索拡張生成 (RAG) パイプライン

###2.1. PubMed 抄録の索引付け

from sentence_transformers import SentenceTransformer
import numpy as np
import faiss
import json

class PubMedRetriever:
    """
    Vector search trên PubMed abstracts với FAISS.
    Embedding model: PubMedBERT fine-tuned cho retrieval (BioSentVec).
    """
    def __init__(self, model_name: str = "pritamdeka/S-PubMedBert-MS-MARCO"):
        self.encoder = SentenceTransformer(model_name)
        self.index = None
        self.documents = []  # List of {"pmid", "title", "abstract"}

    def build_index(self, documents: list[dict]):
        """Build FAISS index từ danh sách abstracts."""
        self.documents = documents
        texts = [f"{d['title']}. {d['abstract']}" for d in documents]

        print(f"Encoding {len(texts)} documents...")
        embeddings = self.encoder.encode(
            texts,
            batch_size=64,
            show_progress_bar=True,
            convert_to_numpy=True,
            normalize_embeddings=True  # Cosine similarity
        )

        # FAISS IndexFlatIP: inner product = cosine similarity (khi normalized)
        dim = embeddings.shape[1]
        self.index = faiss.IndexFlatIP(dim)
        self.index.add(embeddings.astype(np.float32))
        print(f"Index built: {self.index.ntotal} vectors of dim {dim}")

    def retrieve(self, query: str, top_k: int = 5) -> list[dict]:
        """Retrieve top-K most relevant documents."""
        query_embedding = self.encoder.encode(
            [query], normalize_embeddings=True
        ).astype(np.float32)

        scores, indices = self.index.search(query_embedding, top_k)

        results = []
        for score, idx in zip(scores[0], indices[0]):
            doc = self.documents[idx].copy()
            doc["similarity_score"] = round(float(score), 4)
            results.append(doc)
        return results

    def save(self, directory: str):
        os.makedirs(directory, exist_ok=True)
        faiss.write_index(self.index, os.path.join(directory, "faiss.index"))
        with open(os.path.join(directory, "documents.json"), "w") as f:
            json.dump(self.documents, f)

    def load(self, directory: str):
        self.index = faiss.read_index(os.path.join(directory, "faiss.index"))
        with open(os.path.join(directory, "documents.json")) as f:
            self.documents = json.load(f)

###2.2. OpenAI互換APIを備えたRAG

from openai import OpenAI

SYSTEM_PROMPT = """Bạn là một trợ lý y tế AI hỗ trợ bác sĩ tra cứu thông tin.

Nguyên tắc:
1. CHỈ trả lời dựa trên context được cung cấp (các nghiên cứu)
2. LUÔN cite nguồn (PMID hoặc guideline số)
3. Nếu context không đủ, nói rõ "Không tìm thấy bằng chứng đủ mạnh trong cơ sở dữ liệu"
4. KHÔNG đưa ra chẩn đoán cụ thể cho bệnh nhân
5. KHÔNG thay thế tư vấn bác sĩ
6. Với câu hỏi khẩn cấp (đau ngực, khó thở, mất ý thức), luôn khuyên đến cấp cứu ngay"""

class MedicalQAChatbot:
    def __init__(self, retriever: PubMedRetriever, openai_api_key: str):
        self.retriever = retriever
        self.client = OpenAI(api_key=openai_api_key)
        self.conversation_history = []

    def format_context(self, retrieved_docs: list[dict]) -> str:
        context_parts = []
        for i, doc in enumerate(retrieved_docs, 1):
            context_parts.append(
                f"[{i}] PMID: {doc.get('pmid', 'N/A')}\n"
                f"Title: {doc['title']}\n"
                f"Abstract: {doc['abstract'][:500]}...\n"
                f"Similarity: {doc['similarity_score']}"
            )
        return "\n\n".join(context_parts)

    def ask(self, question: str, top_k: int = 5) -> dict:
        """Ask a medical question with RAG."""
        # Safety check TRƯỚC khi retrieve
        safety_check = self._safety_check(question)
        if safety_check["is_emergency"]:
            return {
                "answer": safety_check["emergency_message"],
                "sources": [],
                "is_emergency": True
            }

        # Retrieve relevant documents
        docs = self.retriever.retrieve(question, top_k=top_k)
        context = self.format_context(docs)

        # Build messages
        messages = [
            {"role": "system", "content": SYSTEM_PROMPT},
        ] + self.conversation_history + [
            {
                "role": "user",
                "content": f"""Câu hỏi: {question}

Context từ PubMed (hãy dựa vào đây để trả lời):
{context}

Trả lời dựa trên evidence trên, cite [số thứ tự]:\"\"\""
            }
        ]

        response = self.client.chat.completions.create(
            model="gpt-4o-mini",
            messages=messages,
            temperature=0.1,  # Low temperature cho medical content
            max_tokens=1000
        )

        answer = response.choices[0].message.content

        # Update conversation history (giữ 6 turns gần nhất)
        self.conversation_history.append({"role": "user", "content": question})
        self.conversation_history.append({"role": "assistant", "content": answer})
        if len(self.conversation_history) > 12:
            self.conversation_history = self.conversation_history[-12:]

        return {
            "answer": answer,
            "sources": docs,
            "is_emergency": False,
            "tokens_used": response.usage.total_tokens
        }

    def _safety_check(self, question: str) -> dict:
        """Check câu hỏi có phải emergency không."""
        EMERGENCY_KEYWORDS = [
            "đau ngực dữ dội", "khó thở cấp", "mất ý thức", "co giật",
            "chảy máu không cầm", "ngộ độc", "tự tử", "overdose",
            "chest pain", "can\'t breathe", "unconscious", "seizure"
        ]
        q_lower = question.lower()
        is_emergency = any(kw in q_lower for kw in EMERGENCY_KEYWORDS)
        return {
            "is_emergency": is_emergency,
            "emergency_message": (
                "⚠️ Đây có vẻ là tình huống khẩn cấp y tế. "
                "Hãy gọi ngay 115 (Việt Nam) hoặc đến cơ sở y tế gần nhất NGAY. "
                "Không để chờ AI trả lời trong tình huống khẩn cấp."
            ) if is_emergency else ""
        }

3. 医療ドメイン向けの LLM の微調整

LLM が専門分野についてより詳細な回答を必要とする場合:

from datasets import Dataset
from transformers import (
    AutoModelForCausalLM, AutoTokenizer,
    TrainingArguments, BitsAndBytesConfig
)
from peft import LoraConfig, get_peft_model, TaskType

def setup_medical_lora_finetuning(
    base_model: str = "meta-llama/Llama-3.2-3B-Instruct",
    train_data: list[dict] = None
):
    """
    LoRA fine-tuning cho medical QA.
    
    Tại sao LoRA thay vì full fine-tune?
    - Llama 3B = 3 tỷ parameters × 2 bytes = 6GB VRAM chỉ để load
    - Full fine-tune cần 24-48GB VRAM
    - LoRA: thêm ~1% parameters, cần ~8GB VRAM
    
    Dataset chuẩn: MedQA (USMLE), MedMCQA, PubMedQA
    """
    # 4-bit quantization để fit trong GPU nhỏ
    bnb_config = BitsAndBytesConfig(
        load_in_4bit=True,
        bnb_4bit_use_double_quant=True,
        bnb_4bit_quant_type="nf4",
        bnb_4bit_compute_dtype="bfloat16"
    )

    tokenizer = AutoTokenizer.from_pretrained(base_model)
    model = AutoModelForCausalLM.from_pretrained(
        base_model,
        quantization_config=bnb_config,
        device_map="auto"
    )

    # LoRA config: target attention layers
    lora_config = LoraConfig(
        task_type=TaskType.CAUSAL_LM,
        r=16,           # Rank: higher = more capacity, more memory
        lora_alpha=32,  # Scaling factor
        lora_dropout=0.1,
        target_modules=["q_proj", "v_proj", "k_proj", "o_proj"],
        # Medical specific: cũng target FFN layers
        # target_modules=["q_proj", "v_proj", "k_proj", "o_proj", "gate_proj", "up_proj"]
    )

    model = get_peft_model(model, lora_config)
    model.print_trainable_parameters()
    # Trainable: ~0.7% parameters — hiệu quả cực kỳ cao

    return model, tokenizer

def format_medical_qa_sample(sample: dict) -> str:
    """Format training sample: alpaca style cho medical QA."""
    return (
        f"### Instruction:\n"
        f"Bạn là bác sĩ AI. Hãy trả lời câu hỏi y khoa dựa trên evidence-based medicine.\n\n"
        f"### Question:\n{sample['question']}\n\n"
        f"### Answer:\n{sample['answer']}"
    )

4. 評価: 正確さを超えて

from rouge_score import rouge_scorer

def evaluate_medical_qa(predictions: list[str], references: list[str]) -> dict:
    """
    Metrics cho medical QA:
    - ROUGE-L: text overlap (không đủ cho medical)
    - BERTScore: semantic similarity
    - Medical factuality: check entities against medical KB
    - Citation accuracy: PMID exists & relevant
    """
    scorer = rouge_scorer.RougeScorer(['rougeL'], use_stemmer=True)
    rouge_scores = [
        scorer.score(ref, pred)['rougeL'].fmeasure
        for ref, pred in zip(references, predictions)
    ]

    return {
        "mean_rouge_l": round(sum(rouge_scores) / len(rouge_scores), 4),
        # Thêm BERTScore, medical factuality nếu cần
    }

5. 演習

  1. シンプルな RAG パイプラインを構築します。Entrez API を使用して 1000 件の PubMed 抄録をダウンロードし、FAISS を使用してインデックスを作成し、20 の医療質問でテストします。 Precision@3 を計算します (取得)。

  2. MedQA データセットからの 500 サンプルに対して LoRA を使用して Llama-3.2-3B を微調整します。 USMLE テスト問題のベース モデルと比較します。

  3. 安全ガードレールを実装します。50 のクエリでテストされた 20 の緊急キーワードをリストします。 100% の緊急クエリにフラグが設定されていることを確認します。

レッスン 9: グラフ ニューラル ネットワークを使用した創薬。