已合并
新增 Qwen3 医学场景全参微调与 NPU 推理接入实战教程 #440
gcw_gpcLb7J8创建于 7月28日
新增 Qwen3 医学场景全参微调与 NPU 推理接入实战教程 #440
已合并
共 8 个文件变更+1137-0
| @@ -0,0 +1,75 @@ | |||
| 1 | +# Qwen3 医学场景全参微调与 NPU 推理接入实战 | ||
| 2 | + | ||
| 3 | +在昇腾 NPU 上,参考 [SwanLab Qwen3 医学模型微调教程](https://docs.swanlab.cn/course/llm_train_course/03-sft/4.qwen3-medical-finetune/) 完成 Qwen3-1.7B 的全参微调,训练出一个具备医学问答能力的模型;再把训练产出的权重接入昇腾官方推理框架 [cann-recipes-infer](https://gitcode.com/cann/cann-recipes-infer),在 NPU 上跑起来验证效果。 | ||
| 4 | + | ||
| 5 | +完整代码与分步说明见 `qwen3_medical_sft.ipynb`;本 README 仅覆盖环境相关信息与文件导览。 | ||
| 6 | + | ||
| 7 | +## 体验配置 | ||
| 8 | + | ||
| 9 | +| 项目 | 数值 | | ||
| 10 | +| --- | --- | | ||
| 11 | +| NPU 卡数 | 单卡(本文实测环境为昇腾 Atlas A3,910C 芯片,64GB HBM) | | ||
| 12 | +| 基座模型 | Qwen3-1.7B | | ||
| 13 | +| 训练精度 | BF16 | | ||
| 14 | +| 训练样本数 | 2166 条(训练集)/ 241 条(验证集),共 2407 条,按 9:1 切分 | | ||
| 15 | +| max_length | 2048 | | ||
| 16 | +| 训练步数控制 | `max_steps=680`(约合 5 个 epoch 上限),配合 `EarlyStoppingCallback` 提前停止,实测收敛于约 650 步附近 | | ||
| 17 | +| batch size / 梯度累积 | `per_device_train_batch_size=2`,`gradient_accumulation_steps=8` | | ||
| 18 | +| 推理耗时(`cann-recipes-infer` 实测,两次独立复现结果一致) | Prefill 27.95 ms,Decode 平均 5.18 ms | | ||
| 19 | + | ||
| 20 | +以上配置为本文实际验证环境,未测试更低配置下的可行性;若显存较小,可参考 Notebook 训练部分的 `per_device_train_batch_size`、`gradient_accumulation_steps` 参数,酌情调小 batch size 或开启更激进的显存优化策略。 | ||
| 21 | + | ||
| 22 | +## 前置条件 | ||
| 23 | + | ||
| 24 | +| 项目 | 要求 | | ||
| 25 | +| --- | --- | | ||
| 26 | +| 硬件 | 昇腾 Atlas A2/A3 系列产品(`cann-recipes-infer` 官方支持的产品型号) | | ||
| 27 | +| CANN | 9.0.0 | | ||
| 28 | +| PyTorch / torch_npu | PyTorch 2.7.1 / torch_npu 2.7.1.post4 | | ||
| 29 | +| Python | 3.11 | | ||
| 30 | +| 仓库支持权重 | `cann-recipes-infer` 官方验证过 Qwen3-8B、Qwen2.5-7B-Instruct;本文验证的 Qwen3-1.7B 通过复用 Qwen3-8B 配置模板接入,详见「统一执行器配置模版」相关说明 | | ||
| 31 | + | ||
| 32 | +### 安装教程层依赖 | ||
| 33 | + | ||
| 34 | +在以上底层环境(CANN、PyTorch、torch_npu)已就绪的基础上,还需要安装以下 Python 库: | ||
| 35 | + | ||
| 36 | +```bash | ||
| 37 | +pip install modelscope transformers accelerate swanlab --break-system-packages | ||
| 38 | +``` | ||
| 39 | + | ||
| 40 | +### 配置 SwanLab API Key | ||
| 41 | + | ||
| 42 | +首次使用需要在 [SwanLab 官网](https://swanlab.cn/) 注册账号,在终端环境中完成登录: | ||
| 43 | + | ||
| 44 | +```bash | ||
| 45 | +swanlab login | ||
| 46 | +``` | ||
| 47 | + | ||
| 48 | +执行后按提示粘贴账号 API Key(在 SwanLab 网页端「用户设置 → API Key」页面获取)。 | ||
| 49 | + | ||
| 50 | +### 打开 Notebook | ||
| 51 | + | ||
| 52 | +```bash | ||
| 53 | +jupyter notebook qwen3_medical_sft.ipynb | ||
| 54 | +``` | ||
| 55 | + | ||
| 56 | +## 目录说明 | ||
| 57 | + | ||
| 58 | +| 文件 | 作用 | | ||
| 59 | +| --- | --- | | ||
| 60 | +| `README.md` | 本文档,环境信息与文件导览 | | ||
| 61 | +| `qwen3_medical_sft.ipynb` | 完整代码与分步说明:下载数据 → 加载模型并构建监督标签 → BF16 全参数 SFT → Transformer 原生推理 → cann-recipes-infer 统一执行器部署 | | ||
| 62 | +| `qwen3_medical_sft.yaml` | `cann-recipes-infer` 官方 Qwen3-8B 单卡推理配置模板,接入自定义权重的起点 | | ||
| 63 | +| `prepare_config.py` | 将训练产出的 checkpoint 路径写入运行时 YAML 的 `model_path` 字段 | | ||
| 64 | +| `run_cann_infer.sh` | 从 `models/qwen` 工作目录启动统一执行器,执行推理 | | ||
| 65 | +| `validation.md` | 本次验证的硬件/软件范围、Notebook 完整执行的实测数据、实测遇到并解决的阻塞问题;SwanLab 训练记录与 cann-recipes-infer 推理性能数据见 `qwen3_medical_sft.ipynb` | | ||
| 66 | + | ||
| 67 | +以下文件不在本目录内,属于 `cann-recipes-infer` 仓库本身,会在部署过程中被读取或修改: | ||
| 68 | + | ||
| 69 | +| 路径(相对 cann-recipes-infer 仓库根目录) | 作用 | | ||
| 70 | +| --- | --- | | ||
| 71 | +| `executor/scripts/set_env.sh` | 环境变量配置脚本,需将 `cann_path` 替换为真实 CANN 安装路径 | | ||
| 72 | +| `executor/scripts/infer.sh` | 统一执行器入口脚本,负责加载模型、执行推理 | | ||
| 73 | +| `models/qwen/config/qwen3_8b_1tp.yaml` | 官方提供的 Qwen3-8B 单卡推理配置模板 | | ||
| 74 | +| `dataset/default_prompt.json` | 默认离线推理测试用的 prompt 文件,验证效果时临时替换为领域问题 | | ||
| 75 | + | ||
| @@ -0,0 +1,53 @@ | |||
| 1 | +#!/usr/bin/env python3 | ||
| 2 | +""" | ||
| 3 | +将训练产出的 checkpoint 路径写入统一执行器运行时 YAML 的 model_path 字段。 | ||
| 4 | + | ||
| 5 | +用法: | ||
| 6 | + python3 prepare_config.py <checkpoint绝对路径> [目标yaml路径] | ||
| 7 | + | ||
| 8 | +示例: | ||
| 9 | + python3 prepare_config.py /home/project/qwen3_medical_sft/output_qwen3_medical/final \ | ||
| 10 | + cann-recipes-infer/models/qwen/config/qwen3_custom_1tp.yaml | ||
| 11 | +""" | ||
| 12 | + | ||
| 13 | +import re | ||
| 14 | +import sys | ||
| 15 | +import os | ||
| 16 | + | ||
| 17 | + | ||
| 18 | +def main(): | ||
| 19 | + if len(sys.argv) < 2: | ||
| 20 | + print("错误:请提供 checkpoint 绝对路径作为第一个参数") | ||
| 21 | + print(__doc__) | ||
| 22 | + sys.exit(1) | ||
| 23 | + | ||
| 24 | + checkpoint_path = sys.argv[1] | ||
| 25 | + yaml_path = sys.argv[2] if len(sys.argv) > 2 else \ | ||
| 26 | + "cann-recipes-infer/models/qwen/config/qwen3_custom_1tp.yaml" | ||
| 27 | + | ||
| 28 | + if not os.path.isfile(yaml_path): | ||
| 29 | + print(f"错误:找不到 YAML 文件 {yaml_path}") | ||
| 30 | + print("请先执行: cp cann-recipes-infer/models/qwen/config/qwen3_8b_1tp.yaml " | ||
| 31 | + f"{yaml_path}") | ||
| 32 | + sys.exit(1) | ||
| 33 | + | ||
| 34 | + with open(yaml_path, "r", encoding="utf-8") as f: | ||
| 35 | + content = f.read() | ||
| 36 | + | ||
| 37 | + new_content = re.sub( | ||
| 38 | + r'model_path:\s*".*"', | ||
| 39 | + f'model_path: "{checkpoint_path}"', | ||
| 40 | + content, | ||
| 41 | + ) | ||
| 42 | + | ||
| 43 | + with open(yaml_path, "w", encoding="utf-8") as f: | ||
| 44 | + f.write(new_content) | ||
| 45 | + | ||
| 46 | + print(f"已将 {yaml_path} 中的 model_path 写入: {checkpoint_path}") | ||
| 47 | + for line in new_content.splitlines(): | ||
| 48 | + if "model_path" in line: | ||
| 49 | + print(line.strip()) | ||
| 50 | + | ||
| 51 | + | ||
| 52 | +if __name__ == "__main__": | ||
| 53 | + main() | ||
| @@ -0,0 +1,920 @@ | |||
| 1 | +{ | ||
| 2 | + "cells": [ | ||
| 3 | + { | ||
| 4 | + "cell_type": "markdown", | ||
| 5 | + "metadata": {}, | ||
| 6 | + "source": [ | ||
| 7 | + "# Qwen3 医学场景全参微调与 NPU 推理接入 —— 完整交互体验\n", | ||
| 8 | + "\n", | ||
| 9 | + "本 Notebook 完整演示从数据准备到最终接入 `cann-recipes-infer` 推理部署的全流程,代码均为昇腾 NPU 适配版。\n", | ||
| 10 | + "\n", | ||
| 11 | + "流程分为 5 个部分:\n", | ||
| 12 | + "1. 下载数据\n", | ||
| 13 | + "2. 加载模型并构建监督标签\n", | ||
| 14 | + "3. BF16 全参数 SFT\n", | ||
| 15 | + "4. Transformer 原生推理\n", | ||
| 16 | + "5. cann-recipes-infer 统一执行器部署\n", | ||
| 17 | + "\n", | ||
| 18 | + "> 运行前请确认已完成环境准备(见 `README.md` 中的「前置条件」「体验配置」),本 Notebook 不重复环境安装步骤。\n" | ||
| 19 | + ] | ||
| 20 | + }, | ||
| 21 | + { | ||
| 22 | + "cell_type": "markdown", | ||
| 23 | + "metadata": {}, | ||
| 24 | + "source": [ | ||
| 25 | + "## 1. 下载数据\n", | ||
| 26 | + "\n", | ||
| 27 | + "使用的是 [delicate_medical_r1_data](https://modelscope.cn/datasets/krisfu/delicate_medical_r1_data) 数据集,只取其中 `question`、`think`、`answer` 三列:\n", | ||
| 28 | + "\n", | ||
| 29 | + "- `question`:用户提出的问题,即模型的输入\n", | ||
| 30 | + "- `think`:模型的思考过程\n", | ||
| 31 | + "- `answer`:模型思考完成后的回复内容\n", | ||
| 32 | + "\n", | ||
| 33 | + "下载并按 9:1 划分为训练集(`train.jsonl`)与验证集(`val.jsonl`):\n" | ||
| 34 | + ] | ||
| 35 | + }, | ||
| 36 | + { | ||
| 37 | + "cell_type": "code", | ||
| 38 | + "execution_count": null, | ||
| 39 | + "metadata": {}, | ||
| 40 | + "outputs": [], | ||
| 41 | + "source": [ | ||
| 42 | + "import json\n", | ||
| 43 | + "import random\n", | ||
| 44 | + "import subprocess\n", | ||
| 45 | + "import os\n", | ||
| 46 | + "\n", | ||
| 47 | + "if not os.path.exists(\"delicate_medical_r1_data\"):\n", | ||
| 48 | + " subprocess.run([\n", | ||
| 49 | + " \"git\", \"clone\",\n", | ||
| 50 | + " \"https://www.modelscope.cn/datasets/krisfu/delicate_medical_r1_data.git\"\n", | ||
| 51 | + " ], check=True)\n", | ||
| 52 | + "\n", | ||
| 53 | + "dataset_file = \"delicate_medical_r1_data/r1_data_example.jsonl\"\n", | ||
| 54 | + "print(f\"读取数据集文件: {dataset_file}\")\n", | ||
| 55 | + "\n", | ||
| 56 | + "data = []\n", | ||
| 57 | + "with open(dataset_file, \"r\", encoding=\"utf-8\") as f:\n", | ||
| 58 | + " for line in f:\n", | ||
| 59 | + " if line.strip():\n", | ||
| 60 | + " data.append(json.loads(line))\n", | ||
| 61 | + "\n", | ||
| 62 | + "print(f\"总样本数: {len(data)}\")\n", | ||
| 63 | + "\n", | ||
| 64 | + "random.seed(42)\n", | ||
| 65 | + "random.shuffle(data)\n", | ||
| 66 | + "split_idx = int(len(data) * 0.9)\n", | ||
| 67 | + "train_data = data[:split_idx]\n", | ||
| 68 | + "val_data = data[split_idx:]\n", | ||
| 69 | + "\n", | ||
| 70 | + "with open(\"train.jsonl\", \"w\", encoding=\"utf-8\") as f:\n", | ||
| 71 | + " for item in train_data:\n", | ||
| 72 | + " f.write(json.dumps(item, ensure_ascii=False) + \"\\n\")\n", | ||
| 73 | + "\n", | ||
| 74 | + "with open(\"val.jsonl\", \"w\", encoding=\"utf-8\") as f:\n", | ||
| 75 | + " for item in val_data:\n", | ||
| 76 | + " f.write(json.dumps(item, ensure_ascii=False) + \"\\n\")\n", | ||
| 77 | + "\n", | ||
| 78 | + "print(f\"训练集: {len(train_data)} 条 -> train.jsonl\")\n", | ||
| 79 | + "print(f\"验证集: {len(val_data)} 条 -> val.jsonl\")\n" | ||
| 80 | + ] | ||
| 81 | + }, | ||
| 82 | + { | ||
| 83 | + "cell_type": "markdown", | ||
| 84 | + "metadata": {}, | ||
| 85 | + "source": [ | ||
| 86 | + "## 2. 加载模型并构建监督标签\n", | ||
| 87 | + "\n", | ||
| 88 | + "下载 Qwen3-1.7B 基座模型,并加载到昇腾 NPU 上:\n" | ||
| 89 | + ] | ||
| 90 | + }, | ||
| 91 | + { | ||
| 92 | + "cell_type": "code", | ||
| 93 | + "execution_count": null, | ||
| 94 | + "metadata": {}, | ||
| 95 | + "outputs": [], | ||
| 96 | + "source": [ | ||
| 97 | + "from modelscope import snapshot_download\n", | ||
| 98 | + "\n", | ||
| 99 | + "model_dir = snapshot_download(\"Qwen/Qwen3-1.7B\")\n", | ||
| 100 | + "print(f\"模型下载到: {model_dir}\")\n" | ||
| 101 | + ] | ||
| 102 | + }, | ||
| 103 | + { | ||
| 104 | + "cell_type": "code", | ||
| 105 | + "execution_count": null, | ||
| 106 | + "metadata": {}, | ||
| 107 | + "outputs": [], | ||
| 108 | + "source": [ | ||
| 109 | + "import os\n", | ||
| 110 | + "os.environ[\"HF_HUB_OFFLINE\"] = \"1\"\n", | ||
| 111 | + "os.environ[\"TRANSFORMERS_OFFLINE\"] = \"1\"\n", | ||
| 112 | + "import torch\n", | ||
| 113 | + "import torch_npu\n", | ||
| 114 | + "from transformers import AutoTokenizer, AutoModelForCausalLM\n", | ||
| 115 | + "\n", | ||
| 116 | + "MODEL_PATH = model_dir # 上一步下载得到的路径\n", | ||
| 117 | + "\n", | ||
| 118 | + "tokenizer = AutoTokenizer.from_pretrained(MODEL_PATH, trust_remote_code=True, padding_side=\"right\")\n", | ||
| 119 | + "if tokenizer.pad_token is None:\n", | ||
| 120 | + " tokenizer.pad_token = tokenizer.eos_token\n", | ||
| 121 | + "\n", | ||
| 122 | + "model = AutoModelForCausalLM.from_pretrained(\n", | ||
| 123 | + " MODEL_PATH, trust_remote_code=True, torch_dtype=torch.bfloat16\n", | ||
| 124 | + ").to(\"npu\")\n", | ||
| 125 | + "\n", | ||
| 126 | + "print(f\"NPU 可用: {torch.npu.is_available()}, 设备数: {torch.npu.device_count()}\")\n", | ||
| 127 | + "total_params = sum(p.numel() for p in model.parameters())\n", | ||
| 128 | + "print(f\"总参数量: {total_params / 1e9:.2f}B\")\n" | ||
| 129 | + ] | ||
| 130 | + }, | ||
| 131 | + { | ||
| 132 | + "cell_type": "markdown", | ||
| 133 | + "metadata": {}, | ||
| 134 | + "source": [ | ||
| 135 | + "### 微调前效果(Baseline)\n", | ||
| 136 | + "\n", | ||
| 137 | + "在开始微调之前,先用同样的测试问题跑一遍未经训练的基座模型,作为微调前的效果基线,方便和「4. Transformer 原生推理」小节里微调后的输出直接对比:\n" | ||
| 138 | + ] | ||
| 139 | + }, | ||
| 140 | + { | ||
| 141 | + "cell_type": "code", | ||
| 142 | + "execution_count": null, | ||
| 143 | + "metadata": {}, | ||
| 144 | + "outputs": [], | ||
| 145 | + "source": [ | ||
| 146 | + "SYSTEM_PROMPT = \"你是一个专业的医学推理助手。请根据用户提供的医学问题,进行深入的思考和推理,然后给出专业、准确的回答。\"\n", | ||
| 147 | + "\n", | ||
| 148 | + "test_questions = [\n", | ||
| 149 | + " \"什么是高血压?它的诊断标准是什么?\",\n", | ||
| 150 | + " \"阿司匹林的主要药理作用和常见不良反应有哪些?\",\n", | ||
| 151 | + " \"糖尿病患者出现低血糖反应时应如何紧急处理?\",\n", | ||
| 152 | + "]\n", | ||
| 153 | + "\n", | ||
| 154 | + "model.eval()\n", | ||
| 155 | + "for i, question in enumerate(test_questions):\n", | ||
| 156 | + " print(f\"\\n{'='*60}\")\n", | ||
| 157 | + " print(f\"问题 {i+1}: {question}\")\n", | ||
| 158 | + " print(\"-\" * 60)\n", | ||
| 159 | + "\n", | ||
| 160 | + " messages = [\n", | ||
| 161 | + " {\"role\": \"system\", \"content\": SYSTEM_PROMPT},\n", | ||
| 162 | + " {\"role\": \"user\", \"content\": question},\n", | ||
| 163 | + " ]\n", | ||
| 164 | + " text = tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)\n", | ||
| 165 | + " inputs = tokenizer(text, return_tensors=\"pt\").to(\"npu\")\n", | ||
| 166 | + "\n", | ||
| 167 | + " with torch.no_grad():\n", | ||
| 168 | + " outputs = model.generate(**inputs, max_new_tokens=1536, temperature=0.7, top_p=0.9, do_sample=True)\n", | ||
| 169 | + "\n", | ||
| 170 | + " generated_ids = outputs[0][inputs[\"input_ids\"].shape[1]:]\n", | ||
| 171 | + " response = tokenizer.decode(generated_ids, skip_special_tokens=True)\n", | ||
| 172 | + " print(response)\n" | ||
| 173 | + ] | ||
| 174 | + }, | ||
| 175 | + { | ||
| 176 | + "cell_type": "markdown", | ||
| 177 | + "metadata": {}, | ||
| 178 | + "source": [ | ||
| 179 | + "运行结果示例(微调前):\n", | ||
| 180 | + "\n", | ||
| 181 | + "**问题 1:什么是高血压?它的诊断标准是什么?**\n", | ||
| 182 | + "```\n", | ||
| 183 | + "<think>\n", | ||
| 184 | + "好的,用户问的是高血压的定义和诊断标准。首先,我需要确认高血压的基本概念,包括原发性高血压和继发性高血压的区别。然后,要回顾诊断标准,比如WHO的标准和中国指南的不同之处。还要考虑不同血压值的分类,比如正常高值、高血压前期和确诊高血压。可能用户需要了解诊断的流程,比如测量方法和注意事项。另外,是否需要提到危险因素或并发症?不过问题只问了诊断标准,所以可能不需要深入。要确保信息准确,引用权威来源,比如WHO和中国高血压指南。还要注意用户可能的背景,可能是普通患者或医学生,所以语言要通俗易懂,但保持专业性。最后检查是否有遗漏的关键点,比如血压测量的时机、持续时间等。\n", | ||
| 185 | + "</think>\n", | ||
| 186 | + "\n", | ||
| 187 | + "高血压是指血压持续升高导致器官供血不足和器官损伤的病理状态。根据世界卫生组织(WHO)和中国《高血压防治指南》等权威标准,高血压的诊断需满足以下条件:\n", | ||
| 188 | + "\n", | ||
| 189 | + "---\n", | ||
| 190 | + "\n", | ||
| 191 | + "### 一、高血压的定义\n", | ||
| 192 | + "1. 原发性高血压(占90%以上):病因不明,与遗传、年龄、生活方式等因素相关。\n", | ||
| 193 | + "2. 继发性高血压:由其他疾病(如肾动脉狭窄、嗜铬细胞瘤、慢性肾脏病等)引起。\n", | ||
| 194 | + "\n", | ||
| 195 | + "---\n", | ||
| 196 | + "\n", | ||
| 197 | + "### 二、诊断标准\n", | ||
| 198 | + "#### 1. 诊室血压测量(标准方法)\n", | ||
| 199 | + "- 收缩压(SBP)≥140 mmHg 且 舒张压(DBP)≥90 mmHg,同时满足以下任一条件:\n", | ||
| 200 | + " - 3次以上测量结果符合上述标准。\n", | ||
| 201 | + " - 有靶器官损伤(如左心室肥厚、脑卒中史、蛋白尿等)。\n", | ||
| 202 | + "\n", | ||
| 203 | + "#### 2. 中国高血压指南(2018)\n", | ||
| 204 | + "- 正常高值:SBP 120-139 mmHg 或 DBP 80-89 mmHg。\n", | ||
| 205 | + "- 高血压前期:SBP 140-139 mmHg 或 DBP 90-89 mmHg。\n", | ||
| 206 | + "- 确诊高血压:SBP ≥140 mmHg 且 DBP ≥90 mmHg(需排除白大衣高血压、运动后血压升高等因素)。\n", | ||
| 207 | + "\n", | ||
| 208 | + "---\n", | ||
| 209 | + "\n", | ||
| 210 | + "### 三、诊断流程\n", | ||
| 211 | + "1. 测量时机:晨起静息状态(避免情绪波动、咖啡因、酒精等影响)。\n", | ||
| 212 | + "2. 测量方法:使用经量血压计(水银血压计或电子血压计),确保袖带宽度适中(上臂长度为袖带宽度的2倍)。\n", | ||
| 213 | + "3. 重复测量:至少2次,间隔10分钟以上,取平均值。\n", | ||
| 214 | + "\n", | ||
| 215 | + "---\n", | ||
| 216 | + "\n", | ||
| 217 | + "### 四、特殊情况\n", | ||
| 218 | + "- 白大衣高血压:诊室血压升高但运动后血压正常,需进一步检查。\n", | ||
| 219 | + "- 隐匿性高血压:静息血压升高但运动后正常,需结合临床评估。\n", | ||
| 220 | + "- 妊娠高血压:需区分妊娠期高血压和慢性高血压。\n", | ||
| 221 | + "\n", | ||
| 222 | + "---\n", | ||
| 223 | + "\n", | ||
| 224 | + "### 五、诊断要点\n", | ||
| 225 | + "- 靶器官损伤:如左心室肥厚、脑卒中史、蛋白尿、肾功能异常等。\n", | ||
| 226 | + "- 危险因素:肥胖、高盐饮食、吸烟、酗酒、家族史等。\n", | ||
| 227 | + "\n", | ||
| 228 | + "---\n", | ||
| 229 | + "\n", | ||
| 230 | + "### 总结\n", | ||
| 231 | + "高血压的诊断需结合血压值、靶器官损害及危险因素综合判断。早期发现和干预可显著降低心血管事件风险。若需进一步评估,建议结合临床检查(如心电图、尿常规、肾脏功能等)和动态血压监测。\n", | ||
| 232 | + "```\n", | ||
| 233 | + "\n", | ||
| 234 | + "**问题 2:阿司匹林的主要药理作用和常见不良反应有哪些?**\n", | ||
| 235 | + "```\n", | ||
| 236 | + "<think>\n", | ||
| 237 | + "好的,我现在需要回答用户关于阿司匹林的主要药理作用和常见不良反应的问题。首先,我得回忆一下阿司匹林的基本信息。阿司匹林,也就是水杨酸钠,是一种非甾体抗炎药(NSAIDs),主要用于缓解疼痛、降低体温、减少炎症和退烧。它还有抗血小板的作用,帮助预防血栓形成,因此常用于心脏病和中风的预防。\n", | ||
| 238 | + "\n", | ||
| 239 | + "接下来是药理作用部分。阿司匹林的药理作用主要在于抑制环氧化酶(COX)的活性,特别是COX-1和COX-2。COX-1在胃肠道中起作用,保护胃黏膜,而COX-2则参与炎症反应。抑制COX可以减少前列腺素的合成,从而达到解热、镇痛和抗炎的效果。此外,阿司匹林还具有抗血小板聚集的作用,通过抑制血小板的聚集来预防血栓,这在心血管疾病预防中很重要。\n", | ||
| 240 | + "\n", | ||
| 241 | + "然后是常见不良反应。首先,胃肠道不适是常见的,比如胃痛、胃溃疡、消化性溃疡,这可能是因为COX-1抑制导致胃黏膜保护作用减弱。其次,出血风险,特别是消化道出血和脑出血,因为抗血小板作用可能增加出血风险。还有过敏反应,如皮疹、荨麻疹,严重时可能引起过敏性休克。另外,长期使用可能引起肝肾功能异常,或者水杨酸反应,表现为恶心、呕吐、头痛等。还有可能引起中枢神经系统副作用,比如嗜睡、头晕,或者在高剂量下引起中枢抑制。\n", | ||
| 242 | + "\n", | ||
| 243 | + "还需要考虑其他不良反应,比如肾功能影响,长期使用可能导致肾功能下降,尤其是老年人或有肾病史的患者。此外,阿司匹林可能引起心血管方面的副作用,如心律失常,尤其是在高剂量或长期使用的情况下。\n", | ||
| 244 | + "\n", | ||
| 245 | + "不过,我需要确认这些信息是否准确。比如,是否所有不良反应都正确?比如,水杨酸反应是否确实存在,以及它的表现。还有,是否有其他常见的不良反应,比如心血管事件,或者与其他药物的相互作用?\n", | ||
| 246 | + "\n", | ||
| 247 | + "另外,需要注意阿司匹林的剂量和使用方法,比如长期使用可能有胃肠道副作用,而短期使用可能风险较低。还有,不同人群的反应可能不同,比如老年人或有胃病史的人更容易出现胃肠道问题。\n", | ||
| 248 | + "\n", | ||
| 249 | + "总结起来,药理作用部分要强调COX抑制和抗血小板作用,不良反应则包括胃肠道、出血、过敏、肝肾功能、中枢神经系统副作用等。需要分点列出,确保准确性和全面性。\n", | ||
| 250 | + "</think>\n", | ||
| 251 | + "\n", | ||
| 252 | + "阿司匹林(乙酰水杨酸)是一种常见的非甾体抗炎药(NSAIDs),其主要药理作用和不良反应如下:\n", | ||
| 253 | + "\n", | ||
| 254 | + "---\n", | ||
| 255 | + "\n", | ||
| 256 | + "### 一、主要药理作用\n", | ||
| 257 | + "1. 解热镇痛:通过抑制中枢神经系统前列腺素(PGI2)的合成,减少发热和疼痛感受器的敏感性,达到退热、镇痛效果。\n", | ||
| 258 | + "2. 抗炎作用:抑制炎症反应中的COX-1和COX-2酶活性,减少前列腺素的生成,从而减轻炎症反应。\n", | ||
| 259 | + "3. 抗血小板聚集:通过抑制血小板中TXA2的合成,减少血小板聚集,预防血栓形成,常用于心血管疾病的预防(如心肌梗死、中风)。\n", | ||
| 260 | + "4. 抗风湿作用:对类风湿性关节炎等自身免疫性疾病有一定疗效,但需注意长期使用可能引发的副作用。\n", | ||
| 261 | + "\n", | ||
| 262 | + "---\n", | ||
| 263 | + "\n", | ||
| 264 | + "### 二、常见不良反应\n", | ||
| 265 | + "1. 胃肠道反应:胃肠道溃疡/出血(COX-1抑制导致胃黏膜保护作用减弱,增加胃炎、胃溃疡、消化道出血风险,尤其长期大剂量使用);恶心、呕吐(可能因药物刺激胃部或中枢神经系统抑制作用)。\n", | ||
| 266 | + "2. 出血风险:消化道出血(胃肠道黏膜损伤导致出血,严重时可致呕血、黑便);颅内出血(抗血小板作用可能增加脑出血风险,尤其与肝素等抗凝药合用时);其他出血(如牙龈出血、皮肤瘀斑等)。\n", | ||
| 267 | + "3. 过敏反应:皮疹、荨麻疹(过敏体质者可能出现过敏反应);严重过敏性休克(罕见但可能因药物过敏反应引发)。\n", | ||
| 268 | + "4. 肝肾功能影响:肝功能异常(长期大剂量使用可能引发肝酶升高或肝损伤);肾功能异常(长期使用可能增加肾结石或肾功能下降风险,尤其老年人或肾功能不全者)。\n", | ||
| 269 | + "5. 中枢神经系统副作用:嗜睡、头晕(因药物对中枢神经系统的抑制作用);中枢抑制(高剂量或长期使用可能引起嗜睡、意识模糊等)。\n", | ||
| 270 | + "6. 其他:水杨酸反应(罕见但可能表现为头痛、恶心、呕吐等,多见于长期大剂量使用或与某些药物如利尿剂、抗抑郁药联用时);心血管风险(长期使用可能增加心肌梗死、中风风险,尤其与阿司匹林合用其他抗血小板药物时)。\n", | ||
| 271 | + "\n", | ||
| 272 | + "---\n", | ||
| 273 | + "\n", | ||
| 274 | + "### 三、注意事项\n", | ||
| 275 | + "- 禁忌症:严重肝肾疾病、出血倾向、过敏史、孕妇(妊娠期禁用)、儿童等。\n", | ||
| 276 | + "- 用药原则:需遵医嘱,避免长期大剂量使用,尤其是有胃肠道病史者。\n", | ||
| 277 | + "- 相互作用:与抗凝药(如华法林)、抗血小板药(如氯吡格雷)联用可能增加出血风险,需谨慎。\n", | ||
| 278 | + "\n", | ||
| 279 | + "---\n", | ||
| 280 | + "\n", | ||
| 281 | + "### 总结\n", | ||
| 282 | + "阿司匹林通过抑制COX酶发挥解热、镇痛、抗炎及抗血小板作用,但需警惕胃肠道出血、出血性中风、肝肾损伤等不良反应。合理用药需结合患者个体情况及临床指征。\n", | ||
| 283 | + "```\n", | ||
| 284 | + "\n", | ||
| 285 | + "**问题 3:糖尿病患者出现低血糖反应时应如何紧急处理?**\n", | ||
| 286 | + "```\n", | ||
| 287 | + "<think>\n", | ||
| 288 | + "好的,用户问的是糖尿病患者出现低血糖反应时应如何紧急处理。首先,我需要回忆一下低血糖的定义和症状。低血糖通常是指血糖水平低于正常范围,一般低于3.9 mmol/L。症状可能包括出汗、心慌、颤抖、头晕、恶心、意识模糊,甚至昏迷。\n", | ||
| 289 | + "\n", | ||
| 290 | + "接下来,用户可能想知道具体的处理步骤。首先,应该确认低血糖是否确实发生,然后采取措施。常见的处理方法包括立即摄入快速作用的碳水化合物,比如葡萄糖片、果汁或糖果,然后监测血糖。同时,需要避免进一步的运动,防止血糖进一步下降。\n", | ||
| 291 | + "\n", | ||
| 292 | + "还要考虑患者是否有口服药物,比如胰岛素或其他降糖药,是否需要调整剂量。如果患者有胰岛素,可能需要暂时停用,或者调整使用时间。另外,如果患者有糖尿病并发症,比如酮症酸中毒,可能需要进一步处理,但这种情况通常伴随其他症状,比如恶心、呕吐、腹痛等。\n", | ||
| 293 | + "\n", | ||
| 294 | + "还要注意患者是否有昏迷的情况,这时候可能需要静脉注射葡萄糖,或者在紧急情况下联系医疗人员。同时,要提醒患者以后避免空腹时间过长,规律饮食,避免剧烈运动,以及定期监测血糖。\n", | ||
| 295 | + "\n", | ||
| 296 | + "可能用户是患者或家属,需要明确的步骤和注意事项。需要确保回答准确,步骤清晰,同时提醒可能的并发症和后续措施。比如,如果患者出现昏迷,可能需要立即送医,而不是自行处理。\n", | ||
| 297 | + "\n", | ||
| 298 | + "还要考虑不同类型的糖尿病,比如1型和2型,处理方式是否有差异?不过一般情况下,处理原则相似,但1型糖尿病患者可能需要更频繁的胰岛素调整。不过用户的问题可能更偏向于一般情况,所以保持通用建议即可。\n", | ||
| 299 | + "\n", | ||
| 300 | + "最后,确保回答结构清晰,分点列出步骤,同时加入注意事项,比如避免使用某些药物,监测血糖,后续管理等。需要避免专业术语过多,用通俗易懂的语言解释,同时保持准确性。\n", | ||
| 301 | + "</think>\n", | ||
| 302 | + "\n", | ||
| 303 | + "糖尿病患者出现低血糖反应时,应立即采取以下紧急处理措施,以防止进一步恶化:\n", | ||
| 304 | + "\n", | ||
| 305 | + "---\n", | ||
| 306 | + "\n", | ||
| 307 | + "### 1. 确认低血糖\n", | ||
| 308 | + "- 症状判断:观察患者是否有以下症状(如出汗、心悸、头晕、意识模糊、嗜睡、昏迷等)。\n", | ||
| 309 | + "- 快速评估:确认是否为低血糖(血糖 < 3.9 mmol/L)。\n", | ||
| 310 | + "\n", | ||
| 311 | + "---\n", | ||
| 312 | + "\n", | ||
| 313 | + "### 2. 紧急处理步骤\n", | ||
| 314 | + "#### (1)立即摄入快速作用的碳水化合物\n", | ||
| 315 | + "- 口服途径:葡萄糖片/糖果 15-20克(如含糖的糖果、含糖饮料);果汁/运动饮料 50-100毫升(含糖饮料);含糖食物如饼干、面包、水果(如香蕉、葡萄)。\n", | ||
| 316 | + "- 静脉注射:若患者昏迷或无法口服,立即静脉注射50%葡萄糖液10-20毫升(需由医护人员操作)。\n", | ||
| 317 | + "\n", | ||
| 318 | + "#### (2)暂停胰岛素或降糖药物\n", | ||
| 319 | + "- 胰岛素患者:若使用胰岛素或口服降糖药(如格列美脲、二甲双胍等),需暂停药物,并避免运动。\n", | ||
| 320 | + "- 注意事项:需根据具体药物和病情调整,避免自行调整剂量。\n", | ||
| 321 | + "\n", | ||
| 322 | + "#### (3)监测血糖\n", | ||
| 323 | + "- 立即检测血糖:确认血糖水平是否回升。\n", | ||
| 324 | + "- 持续监测:若患者意识不清或持续低血糖,需持续监测血糖直至恢复。\n", | ||
| 325 | + "\n", | ||
| 326 | + "#### (4)补充水分\n", | ||
| 327 | + "- 避免脱水:低血糖常伴随脱水,需少量多次补充水分(如口服含电解质的饮料)。\n", | ||
| 328 | + "\n", | ||
| 329 | + "---\n", | ||
| 330 | + "\n", | ||
| 331 | + "### 3. 紧急情况处理\n", | ||
| 332 | + "- 昏迷或意识丧失:立即送医,可能需静脉注射葡萄糖或胰高血糖素(需专业医护人员操作)。\n", | ||
| 333 | + "- 酮症酸中毒:若伴随酮体(如恶心、呕吐、腹痛、呼吸深快),需紧急处理酮症酸中毒。\n", | ||
| 334 | + "\n", | ||
| 335 | + "---\n", | ||
| 336 | + "\n", | ||
| 337 | + "### 4. 后续管理\n", | ||
| 338 | + "- 避免空腹过久:规律饮食,避免长时间空腹。\n", | ||
| 339 | + "- 避免剧烈运动:低血糖发作后2小时内避免剧烈活动。\n", | ||
| 340 | + "- 记录血糖变化:记录发作时间、血糖值、用药情况,便于后续调整治疗。\n", | ||
| 341 | + "- 教育患者:告知患者低血糖的识别症状及处理方法,避免再次发生。\n", | ||
| 342 | + "\n", | ||
| 343 | + "---\n", | ||
| 344 | + "\n", | ||
| 345 | + "### 注意事项\n", | ||
| 346 | + "- 避免使用含糖饮料:若患者无法口服,需避免饮用含糖饮料(如可乐、奶茶),以防血糖波动。\n", | ||
| 347 | + "- 特殊人群:如老年人、儿童或有肝肾功能不全者,需谨慎调整处理方式。\n", | ||
| 348 | + "- 药物相互作用:若患者同时使用某些药物(如利尿剂、NSAIDs),需谨慎。\n", | ||
| 349 | + "\n", | ||
| 350 | + "---\n", | ||
| 351 | + "\n", | ||
| 352 | + "### 总结\n", | ||
| 353 | + "低血糖反应的紧急处理核心是快速恢复血糖,优先选择口服碳水化合物或静脉注射葡萄糖,同时避免进一步降糖措施。若症状持续或加重,需立即送医,防止发生昏迷、器官损伤等严重后果。\n", | ||
| 354 | + "```\n", | ||
| 355 | + "\n", | ||
| 356 | + "可以看出,微调前的基座模型已经具备较强的通用医学知识和结构化回答能力(这与 Qwen3-1.7B 本身的预训练能力有关),`<think>` 推理过程也基本完整。微调前后的差异,需结合「4. Transformer 原生推理」小节中同样 3 个问题的微调后回答来对比。\n" | ||
| 357 | + ] | ||
| 358 | + }, | ||
| 359 | + { | ||
| 360 | + "cell_type": "markdown", | ||
| 361 | + "metadata": {}, | ||
| 362 | + "source": [ | ||
| 363 | + "构建监督标签:把 `think` 和 `answer` 拼接成 `<think>...</think>\\n\\n答案` 的格式作为训练目标,并只对 assistant 回复部分计算 loss(prompt 部分标签设为 `-100`,训练时不计入 loss):\n" | ||
| 364 | + ] | ||
| 365 | + }, | ||
| 366 | + { | ||
| 367 | + "cell_type": "code", | ||
| 368 | + "execution_count": null, | ||
| 369 | + "metadata": {}, | ||
| 370 | + "outputs": [], | ||
| 371 | + "source": [ | ||
| 372 | + "import json\n", | ||
| 373 | + "from torch.utils.data import Dataset\n", | ||
| 374 | + "\n", | ||
| 375 | + "SYSTEM_PROMPT = \"你是一个专业的医学推理助手。请根据用户提供的医学问题,进行深入的思考和推理,然后给出专业、准确的回答。\"\n", | ||
| 376 | + "MAX_LENGTH = 2048\n", | ||
| 377 | + "\n", | ||
| 378 | + "\n", | ||
| 379 | + "class MedicalDataset(Dataset):\n", | ||
| 380 | + " def __init__(self, data_path, tokenizer, max_length=MAX_LENGTH):\n", | ||
| 381 | + " super().__init__()\n", | ||
| 382 | + " self.tokenizer = tokenizer\n", | ||
| 383 | + " self.max_length = max_length\n", | ||
| 384 | + " self.data = []\n", | ||
| 385 | + " with open(data_path, \"r\", encoding=\"utf-8\") as f:\n", | ||
| 386 | + " for line in f:\n", | ||
| 387 | + " if line.strip():\n", | ||
| 388 | + " self.data.append(json.loads(line))\n", | ||
| 389 | + " print(f\"加载数据: {data_path}, 样本数: {len(self.data)}\")\n", | ||
| 390 | + "\n", | ||
| 391 | + " def __len__(self):\n", | ||
| 392 | + " return len(self.data)\n", | ||
| 393 | + "\n", | ||
| 394 | + " def __getitem__(self, index):\n", | ||
| 395 | + " item = self.data[index]\n", | ||
| 396 | + " question = item[\"question\"]\n", | ||
| 397 | + " think = item[\"think\"]\n", | ||
| 398 | + " answer = item[\"answer\"]\n", | ||
| 399 | + "\n", | ||
| 400 | + " # 构建监督标签:think + answer 拼接为训练目标\n", | ||
| 401 | + " assistant_content = f\"<think>{think}</think>\\n\\n{answer}\"\n", | ||
| 402 | + "\n", | ||
| 403 | + " messages = [\n", | ||
| 404 | + " {\"role\": \"system\", \"content\": SYSTEM_PROMPT},\n", | ||
| 405 | + " {\"role\": \"user\", \"content\": question},\n", | ||
| 406 | + " {\"role\": \"assistant\", \"content\": assistant_content},\n", | ||
| 407 | + " ]\n", | ||
| 408 | + " full_text = self.tokenizer.apply_chat_template(\n", | ||
| 409 | + " messages, tokenize=False, add_generation_prompt=False,\n", | ||
| 410 | + " )\n", | ||
| 411 | + "\n", | ||
| 412 | + " prompt_messages = [\n", | ||
| 413 | + " {\"role\": \"system\", \"content\": SYSTEM_PROMPT},\n", | ||
| 414 | + " {\"role\": \"user\", \"content\": question},\n", | ||
| 415 | + " ]\n", | ||
| 416 | + " prompt_text = self.tokenizer.apply_chat_template(\n", | ||
| 417 | + " prompt_messages, tokenize=False, add_generation_prompt=True,\n", | ||
| 418 | + " )\n", | ||
| 419 | + "\n", | ||
| 420 | + " full_ids = self.tokenizer(full_text, max_length=self.max_length, truncation=True, padding=False)\n", | ||
| 421 | + " prompt_ids = self.tokenizer(prompt_text, max_length=self.max_length, truncation=True, padding=False)\n", | ||
| 422 | + "\n", | ||
| 423 | + " input_ids = full_ids[\"input_ids\"]\n", | ||
| 424 | + " attention_mask = full_ids[\"attention_mask\"]\n", | ||
| 425 | + " prompt_len = len(prompt_ids[\"input_ids\"])\n", | ||
| 426 | + " # prompt 部分标签设为 -100,训练时不计入 loss,只监督 assistant 回复部分\n", | ||
| 427 | + " labels = [-100] * prompt_len + input_ids[prompt_len:]\n", | ||
| 428 | + "\n", | ||
| 429 | + " if len(labels) != len(input_ids):\n", | ||
| 430 | + " labels = labels[:len(input_ids)]\n", | ||
| 431 | + "\n", | ||
| 432 | + " return {\"input_ids\": input_ids, \"attention_mask\": attention_mask, \"labels\": labels}\n", | ||
| 433 | + "\n", | ||
| 434 | + "\n", | ||
| 435 | + "train_dataset = MedicalDataset(\"train.jsonl\", tokenizer)\n", | ||
| 436 | + "val_dataset = MedicalDataset(\"val.jsonl\", tokenizer)\n" | ||
| 437 | + ] | ||
| 438 | + }, | ||
| 439 | + { | ||
| 440 | + "cell_type": "markdown", | ||
| 441 | + "metadata": {}, | ||
| 442 | + "source": [ | ||
| 443 | + "## 3. BF16 全参数 SFT\n", | ||
| 444 | + "\n", | ||
| 445 | + "使用 BF16 精度进行全参数监督微调(不使用 LoRA),并通过 `report_to=\"swanlab\"` 记录训练过程:\n" | ||
| 446 | + ] | ||
| 447 | + }, | ||
| 448 | + { | ||
| 449 | + "cell_type": "code", | ||
| 450 | + "execution_count": null, | ||
| 451 | + "metadata": {}, | ||
| 452 | + "outputs": [], | ||
| 453 | + "source": [ | ||
| 454 | + "import swanlab\n", | ||
| 455 | + "from transformers import DataCollatorForSeq2Seq, TrainingArguments, Trainer, EarlyStoppingCallback\n", | ||
| 456 | + "\n", | ||
| 457 | + "OUTPUT_DIR = \"./output_qwen3_medical\"\n", | ||
| 458 | + "\n", | ||
| 459 | + "swanlab.init(project=\"qwen3-medical-npu-finetune\", experiment_name=\"qwen3-1.7b-medical-npu\")\n", | ||
| 460 | + "\n", | ||
| 461 | + "data_collator = DataCollatorForSeq2Seq(\n", | ||
| 462 | + " tokenizer=tokenizer, padding=True, max_length=MAX_LENGTH, label_pad_token_id=-100,\n", | ||
| 463 | + ")\n", | ||
| 464 | + "\n", | ||
| 465 | + "training_args = TrainingArguments(\n", | ||
| 466 | + " output_dir=OUTPUT_DIR,\n", | ||
| 467 | + " per_device_train_batch_size=2,\n", | ||
| 468 | + " per_device_eval_batch_size=2,\n", | ||
| 469 | + " gradient_accumulation_steps=8,\n", | ||
| 470 | + " learning_rate=2e-5,\n", | ||
| 471 | + " num_train_epochs=1,\n", | ||
| 472 | + " max_steps=680, # 上限设为约 5 个 epoch,配合早停实际会更早停止\n", | ||
| 473 | + " warmup_ratio=0.05,\n", | ||
| 474 | + " bf16=True,\n", | ||
| 475 | + " logging_steps=10,\n", | ||
| 476 | + " save_steps=200,\n", | ||
| 477 | + " eval_strategy=\"steps\",\n", | ||
| 478 | + " eval_steps=20,\n", | ||
| 479 | + " save_total_limit=3,\n", | ||
| 480 | + " load_best_model_at_end=True,\n", | ||
| 481 | + " metric_for_best_model=\"eval_loss\",\n", | ||
| 482 | + " greater_is_better=False,\n", | ||
| 483 | + " report_to=\"swanlab\",\n", | ||
| 484 | + " dataloader_pin_memory=False,\n", | ||
| 485 | + " remove_unused_columns=False,\n", | ||
| 486 | + " gradient_checkpointing=True,\n", | ||
| 487 | + " gradient_checkpointing_kwargs={\"use_reentrant\": False},\n", | ||
| 488 | + ")\n", | ||
| 489 | + "\n", | ||
| 490 | + "trainer = Trainer(\n", | ||
| 491 | + " model=model,\n", | ||
| 492 | + " args=training_args,\n", | ||
| 493 | + " train_dataset=train_dataset,\n", | ||
| 494 | + " eval_dataset=val_dataset,\n", | ||
| 495 | + " data_collator=data_collator,\n", | ||
| 496 | + " callbacks=[EarlyStoppingCallback(early_stopping_patience=5)], # 连续5次eval无改善则自动停止\n", | ||
| 497 | + ")\n", | ||
| 498 | + "\n", | ||
| 499 | + "print(\"开始训练...\")\n", | ||
| 500 | + "trainer.train()\n", | ||
| 501 | + "\n", | ||
| 502 | + "print(\"保存模型...\")\n", | ||
| 503 | + "import os\n", | ||
| 504 | + "final_dir = os.path.join(OUTPUT_DIR, \"final\")\n", | ||
| 505 | + "trainer.save_model(final_dir)\n", | ||
| 506 | + "tokenizer.save_pretrained(final_dir)\n", | ||
| 507 | + "print(f\"模型已保存到: {final_dir}\")\n" | ||
| 508 | + ] | ||
| 509 | + }, | ||
| 510 | + { | ||
| 511 | + "cell_type": "markdown", | ||
| 512 | + "metadata": {}, | ||
| 513 | + "source": [ | ||
| 514 | + "训练结束后,取验证集前 3 条数据生成回复,记录到 SwanLab 的 `Prediction` 中,便于在网页端查看训练效果:\n" | ||
| 515 | + ] | ||
| 516 | + }, | ||
| 517 | + { | ||
| 518 | + "cell_type": "code", | ||
| 519 | + "execution_count": null, | ||
| 520 | + "metadata": {}, | ||
| 521 | + "outputs": [], | ||
| 522 | + "source": [ | ||
| 523 | + "model.eval()\n", | ||
| 524 | + "sample_data = []\n", | ||
| 525 | + "with open(\"val.jsonl\", \"r\", encoding=\"utf-8\") as f:\n", | ||
| 526 | + " for i, line in enumerate(f):\n", | ||
| 527 | + " if i >= 3:\n", | ||
| 528 | + " break\n", | ||
| 529 | + " sample_data.append(json.loads(line))\n", | ||
| 530 | + "\n", | ||
| 531 | + "prediction_texts = []\n", | ||
| 532 | + "for item in sample_data:\n", | ||
| 533 | + " question = item[\"question\"]\n", | ||
| 534 | + " messages = [\n", | ||
| 535 | + " {\"role\": \"system\", \"content\": SYSTEM_PROMPT},\n", | ||
| 536 | + " {\"role\": \"user\", \"content\": question},\n", | ||
| 537 | + " ]\n", | ||
| 538 | + " text = tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)\n", | ||
| 539 | + " inputs = tokenizer(text, return_tensors=\"pt\").to(\"npu\")\n", | ||
| 540 | + " with torch.no_grad():\n", | ||
| 541 | + " outputs = model.generate(**inputs, max_new_tokens=1024)\n", | ||
| 542 | + " generated_ids = outputs[0][inputs[\"input_ids\"].shape[1]:]\n", | ||
| 543 | + " response = tokenizer.decode(generated_ids, skip_special_tokens=True)\n", | ||
| 544 | + " print(f\"Question: {question}\\n\\nLLM: {response}\\n\")\n", | ||
| 545 | + " prediction_texts.append(swanlab.Text(f\"Question: {question}\\n\\nLLM: {response}\"))\n", | ||
| 546 | + "\n", | ||
| 547 | + "swanlab.log({\"Prediction\": prediction_texts})\n", | ||
| 548 | + "swanlab.finish()\n" | ||
| 549 | + ] | ||
| 550 | + }, | ||
| 551 | + { | ||
| 552 | + "cell_type": "markdown", | ||
| 553 | + "metadata": {}, | ||
| 554 | + "source": [ | ||
| 555 | + "### SwanLab 云端记录\n", | ||
| 556 | + "\n", | ||
| 557 | + "训练过程与 `Prediction` 样例记录于 SwanLab:\n", | ||
| 558 | + "\n", | ||
| 559 | + "- 项目名:`qwen3-medical-npu-finetune`\n", | ||
| 560 | + "- 实验名:`qwen3-1.7b-medical-npu`\n", | ||
| 561 | + "- 实验链接:https://swanlab.cn/@catherrrrrine/qwen3-medical-npu-finetune\n", | ||
| 562 | + "\n", | ||
| 563 | + "> 该项目已在 SwanLab 设为公开(Public),无需登录 SwanLab 账号即可通过上述链接直接查看训练曲线与 `Prediction` 样例。\n", | ||
| 564 | + "\n", | ||
| 565 | + "`eval/loss` 与 `train/loss` 均随训练步数稳定下降并收敛,全程未见过拟合或发散:\n", | ||
| 566 | + "\n", | ||
| 567 | + "\n", | ||
| 568 | + "\n", | ||
| 569 | + "`eval/loss`:从初始约 1.41 快速下降,200 步左右进入缓慢收敛区间,训练至约 650 步时稳定在 1.16 左右,全程未出现回升,说明未发生过拟合。\n", | ||
| 570 | + "\n", | ||
| 571 | + "\n", | ||
| 572 | + "\n", | ||
| 573 | + "`train/loss`:整体趋势与 `eval/loss` 一致,从初始约 1.93 快速下降,150 步后进入平稳震荡下降区间,训练结束时收敛至约 1.07,与验证集损失走势吻合。\n" | ||
| 574 | + ] | ||
| 575 | + }, | ||
| 576 | + { | ||
| 577 | + "cell_type": "markdown", | ||
| 578 | + "metadata": {}, | ||
| 579 | + "source": [ | ||
| 580 | + "## 4. Transformer 原生推理\n", | ||
| 581 | + "\n", | ||
| 582 | + "在接入 `cann-recipes-infer` 之前,先用原生 `transformers` 单独验证训练产出的 checkpoint 能否正常推理,作为独立于推理框架的基准测试。测试问题与「2. 微调前效果(Baseline)」完全一致,可直接对比微调前后的回答质量:\n" | ||
| 583 | + ] | ||
| 584 | + }, | ||
| 585 | + { | ||
| 586 | + "cell_type": "code", | ||
| 587 | + "execution_count": null, | ||
| 588 | + "metadata": {}, | ||
| 589 | + "outputs": [], | ||
| 590 | + "source": [ | ||
| 591 | + "import torch\n", | ||
| 592 | + "import torch_npu\n", | ||
| 593 | + "from transformers import AutoTokenizer, AutoModelForCausalLM\n", | ||
| 594 | + "\n", | ||
| 595 | + "CHECKPOINT_PATH = \"./output_qwen3_medical/final\"\n", | ||
| 596 | + "SYSTEM_PROMPT = \"你是一个专业的医学推理助手。请根据用户提供的医学问题,进行深入的思考和推理,然后给出专业、准确的回答。\"\n", | ||
| 597 | + "\n", | ||
| 598 | + "test_questions = [\n", | ||
| 599 | + " \"什么是高血压?它的诊断标准是什么?\",\n", | ||
| 600 | + " \"阿司匹林的主要药理作用和常见不良反应有哪些?\",\n", | ||
| 601 | + " \"糖尿病患者出现低血糖反应时应如何紧急处理?\",\n", | ||
| 602 | + "]\n", | ||
| 603 | + "\n", | ||
| 604 | + "infer_tokenizer = AutoTokenizer.from_pretrained(CHECKPOINT_PATH, trust_remote_code=True)\n", | ||
| 605 | + "infer_model = AutoModelForCausalLM.from_pretrained(\n", | ||
| 606 | + " CHECKPOINT_PATH, trust_remote_code=True, torch_dtype=torch.bfloat16\n", | ||
| 607 | + ").to(\"npu\")\n", | ||
| 608 | + "infer_model.eval()\n", | ||
| 609 | + "\n", | ||
| 610 | + "for i, question in enumerate(test_questions):\n", | ||
| 611 | + " print(f\"\\n{'='*60}\")\n", | ||
| 612 | + " print(f\"问题 {i+1}: {question}\")\n", | ||
| 613 | + " print(\"-\" * 60)\n", | ||
| 614 | + "\n", | ||
| 615 | + " messages = [\n", | ||
| 616 | + " {\"role\": \"system\", \"content\": SYSTEM_PROMPT},\n", | ||
| 617 | + " {\"role\": \"user\", \"content\": question},\n", | ||
| 618 | + " ]\n", | ||
| 619 | + " text = infer_tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)\n", | ||
| 620 | + " inputs = infer_tokenizer(text, return_tensors=\"pt\").to(\"npu\")\n", | ||
| 621 | + "\n", | ||
| 622 | + " with torch.no_grad():\n", | ||
| 623 | + " outputs = infer_model.generate(**inputs, max_new_tokens=1536, temperature=0.7, top_p=0.9, do_sample=True)\n", | ||
| 624 | + "\n", | ||
| 625 | + " generated_ids = outputs[0][inputs[\"input_ids\"].shape[1]:]\n", | ||
| 626 | + " response = infer_tokenizer.decode(generated_ids, skip_special_tokens=True)\n", | ||
| 627 | + " print(response)\n" | ||
| 628 | + ] | ||
| 629 | + }, | ||
| 630 | + { | ||
| 631 | + "cell_type": "markdown", | ||
| 632 | + "metadata": {}, | ||
| 633 | + "source": [ | ||
| 634 | + "运行结果示例(微调后):\n", | ||
| 635 | + "\n", | ||
| 636 | + "**问题 1:什么是高血压?它的诊断标准是什么?**\n", | ||
| 637 | + "```\n", | ||
| 638 | + "<think>\n", | ||
| 639 | + "嗯,用户问的是高血压的定义和诊断标准,我需要先回忆一下相关的医学知识。高血压,也就是高血压病,通常是指血压持续升高的情况。首先,我得确定血压的正常范围,因为诊断标准里肯定有这个基础。比如,正常血压是收缩压<120mmHg和舒张压<80mmHg,而高血压就是收缩压≥140mmHg或者舒张压≥90mmHg的情况。\n", | ||
| 640 | + "\n", | ||
| 641 | + "接下来,用户可能想知道诊断的具体标准,不仅仅是数值,还要考虑什么因素?我记得高血压的诊断可能需要结合多个指标,比如持续的血压升高,以及是否有症状。比如,如果患者有高血压症状,或者血压持续升高超过一定时间,可能需要进一步检查,比如24小时动态血压监测或者多次测量。\n", | ||
| 642 | + "\n", | ||
| 643 | + "然后,我需要考虑诊断的分级,比如原发性高血压和继发性高血压的区别,但用户的问题里可能不需要这么详细,但诊断标准中可能包括这些信息。不过根据提供的答案,用户可能只需要基本的数值标准,所以可能不需要深入分级。\n", | ||
| 644 | + "\n", | ||
| 645 | + "另外,诊断标准里提到的\"血压持续升高\"和\"存在症状\"是否准确?比如,有些患者可能没有症状,但血压仍然高,这时候可能需要动态监测。另外,是否需要排除其他疾病导致的高血压?比如药物影响,或者某些病理情况,但诊断标准可能主要基于血压数值和症状。\n", | ||
| 646 | + "\n", | ||
| 647 | + "还要注意单位是否正确,比如mmHg,用户的问题里已经用了这个单位,所以没问题。另外,可能需要提到高血压的分类,比如轻度、中度、重度,但答案里可能没有,所以保持简洁。\n", | ||
| 648 | + "\n", | ||
| 649 | + "有没有可能用户是医学生或者患者,需要更详细的解释?比如,为什么血压超过某个数值,或者高血压的风险因素?不过问题只问定义和标准,所以可能不需要扩展,但答案里已经涵盖了这些内容。\n", | ||
| 650 | + "\n", | ||
| 651 | + "总结一下,答案的结构应该是先定义,然后给出血压数值标准,再说明需要持续升高和症状存在,可能还要提到诊断方法如动态监测。这样用户就能理解了。需要确保没有遗漏关键点,比如诊断的标准不仅仅是数值,还包括症状和时间因素。同时,可能还要提到高血压的类型,但根据提供的答案,可能不需要,所以保持答案的简洁性。\n", | ||
| 652 | + "</think>\n", | ||
| 653 | + "\n", | ||
| 654 | + "高血压,又称高血压病,是指以血管壁的弹性减退和血管硬化为特征的疾病,导致血液流动时对血管壁的侧压增加,即血压升高。根据世界卫生组织的标准,高血压的诊断标准是:如果一个人的血压持续升高,且在安静状态下,收缩压达到或超过140毫米汞柱,或者舒张压达到或超过90毫米汞柱,就可以诊断为高血压。此外,患者还需要存在症状,或者有明确的高血压病史。这种病的诊断通常需要结合患者的血压水平、症状以及进一步的检查来确认。\n", | ||
| 655 | + "```\n", | ||
| 656 | + "\n", | ||
| 657 | + "**问题 2:阿司匹林的主要药理作用和常见不良反应有哪些?**\n", | ||
| 658 | + "```\n", | ||
| 659 | + "<think>\n", | ||
| 660 | + "嗯,用户问的是阿司匹林的主要药理作用和常见不良反应,我需要先回忆一下阿司匹林的基本知识。阿司匹林,也就是水杨酸乙酯,属于非甾体抗炎药(NSAIDs)类,对吧?那它的主要作用应该和抗炎、镇痛、解热有关,对吧?我记得NSAIDs的作用机制是抑制环氧合酶(COX),所以阿司匹林可能通过抑制COX-1和COX-2来减少前列腺素的合成,从而达到这些效果。\n", | ||
| 661 | + "\n", | ||
| 662 | + "接下来是药理作用部分。抗炎方面,确实,阿司匹林可以减轻炎症反应,比如关节炎、肌肉疼痛之类的。镇痛的话,应该也是通过抑制前列腺素,减少疼痛信号传递。解热的话,可能也是同样的机制,降低体温。另外,可能还有抗血小板聚集的作用,对吧?因为阿司匹林常用于预防心脏病发作,所以抗血小板也是它的作用之一。\n", | ||
| 663 | + "\n", | ||
| 664 | + "然后是不良反应。常见的应该包括胃肠道出血、溃疡,因为NSAIDs可能损伤胃肠道黏膜。还有可能引起过敏反应,比如皮疹、哮喘发作。另外,长期使用可能有肝肾毒性,比如肝酶升高、肾功能异常。还有,阿司匹林可能引发水杨酸反应,比如严重的胃肠道反应或者全身性反应,比如皮疹、发热、头痛,严重的话甚至休克。另外,药物相互作用方面,比如与抗凝药(如华法林)合用可能增加出血风险,所以需要提醒用户注意药物相互作用。\n", | ||
| 665 | + "\n", | ||
| 666 | + "用户可能想知道为什么这些不良反应会发生,比如为什么胃肠道出血?可能因为COX-1抑制导致胃黏膜保护作用减弱,而COX-2抑制可能对肠道黏膜影响较小,所以COX-1抑制更主要。另外,过敏反应可能是因为免疫系统对药物成分的反应,或者药物本身引起的。\n", | ||
| 667 | + "\n", | ||
| 668 | + "用户可能没有直接问,但深层需求可能是想了解如何正确使用阿司匹林,或者在用药时需要注意什么。所以回答时不仅要列出药理作用和不良反应,还要强调注意事项,比如饭后服用减少胃刺激,避免与其他抗凝药合用,以及长期使用的风险。\n", | ||
| 669 | + "\n", | ||
| 670 | + "需要确认是否有遗漏的不良反应,比如对肾脏的影响,或者是否还有其他副作用。比如,阿司匹林可能引起心血管事件的风险,比如心脏病发作,但主要还是胃肠道和过敏反应。另外,长期使用可能有其他副作用,比如消化道溃疡、出血,还有可能影响血液凝固,导致出血风险增加。\n", | ||
| 671 | + "\n", | ||
| 672 | + "总结一下,药理作用包括抗炎、镇痛、解热和抗血小板;不良反应包括胃肠道出血、溃疡、过敏反应、肝肾毒性、水杨酸反应、药物相互作用。注意事项要提到避免与其他抗凝药合用,饭后服用,以及长期使用的风险。需要确保信息准确,比如水杨酸反应的具体表现,比如皮疹、发热、头痛,严重时休克。另外,抗血小板作用可能在预防血栓方面有用,但出血风险也存在,所以需要权衡利弊。\n", | ||
| 673 | + "\n", | ||
| 674 | + "可能用户是患者或医学生,需要明确且易懂的解释,所以用词要准确但不过于专业。比如解释COX的作用机制时,可以简单说抑制前列腺素合成,减少炎症和疼痛。不良反应部分要分点说明,但用户给的回复已经分点,可能需要保持结构一致。\n", | ||
| 675 | + "\n", | ||
| 676 | + "最后检查是否有错误,比如阿司匹林的抗血小板作用是否正确,是的,它确实通过抑制血栓烷A2的合成,促进血小板聚集的抑制,所以是抗血小板的。另外,肝肾毒性可能需要提到长期或大剂量使用,或者与其他药物同时使用时的风险增加。水杨酸反应可能需要解释为什么会有,比如药物本身的副作用,或者与某些药物相互作用导致。\n", | ||
| 677 | + "</think>\n", | ||
| 678 | + "\n", | ||
| 679 | + "阿司匹林的主要药理作用包括抗炎、镇痛、解热以及抗血小板的作用。其抗血小板的作用机制是通过抑制血栓烷A2的合成,从而抑制血小板的聚集。此外,阿司匹林还具有抗炎、镇痛、解热等作用,这些作用均来源于它对COX酶的抑制。\n", | ||
| 680 | + "\n", | ||
| 681 | + "阿司匹林的常见不良反应包括胃肠道出血、溃疡、过敏反应、肝肾毒性、水杨酸反应、药物相互作用等。在用药过程中,需要注意避免与其他抗凝药物(如华法林)同时使用,以免增加出血风险。同时,阿司匹林可能引起胃肠道不适,因此建议在饭后服用以减少胃肠道刺激。对于长期或大剂量使用阿司匹林的患者,应密切监测肝肾功能,以预防肝肾毒性。\n", | ||
| 682 | + "```\n", | ||
| 683 | + "\n", | ||
| 684 | + "**问题 3:糖尿病患者出现低血糖反应时应如何紧急处理?**\n", | ||
| 685 | + "```\n", | ||
| 686 | + "<think>\n", | ||
| 687 | + "嗯,用户问的是糖尿病患者出现低血糖反应时的紧急处理方法。首先,我需要回忆一下低血糖的定义和症状。低血糖通常指血糖浓度低于正常范围,可能影响大脑功能,导致头晕、出汗、心悸等症状。然后,我需要考虑糖尿病患者的特殊性,他们可能有胰岛素分泌不足、口服降糖药过量或者运动后血糖波动等问题。\n", | ||
| 688 | + "\n", | ||
| 689 | + "接下来,用户提供的答案已经列出了几个步骤:立即进食含糖食物、补充葡萄糖、监测血糖、再次进食、避免药物影响、及时就医。我需要确保这些步骤是否全面且正确。比如,是否应该建议立即补充含糖食物,而不是直接注射葡萄糖?因为注射葡萄糖可能引起低血糖加重,所以答案中的\"口服葡萄糖\"是正确的。然后,补充葡萄糖后监测血糖是必要的,但用户提供的答案中提到的再次进食是否必要呢?可能在某些情况下,比如患者无法口服,可能需要静脉注射,但答案里没有提到,可能需要确认是否足够。\n", | ||
| 690 | + "\n", | ||
| 691 | + "另外,答案中提到避免药物影响,可能指的是不要立即再次使用胰岛素或其他降糖药,但有时可能需要根据情况调整。比如,如果患者在运动后低血糖,可能需要暂停运动并补充能量。但答案中的建议可能已经足够,所以不需要深入讨论药物调整。\n", | ||
| 692 | + "\n", | ||
| 693 | + "用户可能希望得到一个简明但全面的回答,所以需要确保每个步骤都正确,并且逻辑清晰。可能需要检查是否有遗漏的步骤,比如是否应该建议联系医生或家人,或者是否需要监测血糖直到稳定。此外,是否需要提醒患者下次避免类似情况?\n", | ||
| 694 | + "\n", | ||
| 695 | + "另外,用户的问题中提到\"紧急处理\",所以步骤需要快速有效。比如,立即进食是关键,而补充葡萄糖后再次进食可能用于补充能量,避免低血糖再次发生。同时,监测血糖直到恢复到正常范围,这样能确保处理有效。\n", | ||
| 696 | + "\n", | ||
| 697 | + "可能还需要考虑患者的个体差异,比如是否有胃肠道问题,是否能口服,或者是否需要静脉输注葡萄糖。但答案中的方法已经足够一般性,所以应该没问题。总结下来,答案已经覆盖了主要处理步骤,所以应该没问题。但需要确保用词准确,比如\"含糖食物\"是否包括水果、糖果等,而不仅仅是糖水。可能需要确认是否应该建议含水的糖分,但通常含糖食物是合适的。\n", | ||
| 698 | + "\n", | ||
| 699 | + "最后,检查是否有更好的方法或补充点,比如是否应该建议患者在处理后立即联系医生,或者是否需要进一步的检查。但答案中已经提到及时就医,所以可能已经足够。因此,用户提供的答案是正确的,但需要确保没有错误,比如是否应该先口服葡萄糖再补充食物,或者是否应该先补充食物再注射葡萄糖。根据答案中的步骤,先口服葡萄糖,然后监测,再补充食物,这可能是因为患者可能无法立即口服,或者需要先补充葡萄糖以迅速提升血糖,然后再通过食物补充能量。所以这个顺序是合理的。\n", | ||
| 700 | + "\n", | ||
| 701 | + "因此,综合以上思考,答案中的步骤是正确的,覆盖了紧急处理的关键点,并且逻辑清晰。\n", | ||
| 702 | + "</think>\n", | ||
| 703 | + "\n", | ||
| 704 | + "糖尿病患者出现低血糖反应时,应立即补充含糖食物,如糖果、果汁、面包等,以迅速提高血糖水平。同时,应监测血糖,观察是否有改善。如果血糖仍未恢复正常,应再次补充含糖食物。在处理过程中,应避免使用可能降低血糖的药物,如胰岛素,以防止血糖进一步下降。如果情况严重,应及时就医,以便获得进一步的治疗。\n", | ||
| 705 | + "```\n", | ||
| 706 | + "\n", | ||
| 707 | + "对比「2. 微调前效果(Baseline)」中同样 3 个问题的回答可以看出:微调前的基座模型倾向于用 `###` 分级标题、分点列表把回答组织成结构化的\"文档\"形式;微调后的模型仍保留了完整的 `<think>` 推理过程,但最终回答收敛为更直接、更口语化的连续段落,风格上更贴近医学问答数据集本身的标注格式。\n" | ||
| 708 | + ] | ||
| 709 | + }, | ||
| 710 | + { | ||
| 711 | + "cell_type": "markdown", | ||
| 712 | + "metadata": {}, | ||
| 713 | + "source": [ | ||
| 714 | + "## 5. cann-recipes-infer 统一执行器部署\n", | ||
| 715 | + "\n", | ||
| 716 | + "将训练产出的 checkpoint 接入昇腾官方推理框架 [cann-recipes-infer](https://gitcode.com/cann/cann-recipes-infer),通过统一执行器 `executor/scripts/infer.sh` 完成部署与推理验证。\n", | ||
| 717 | + "\n", | ||
| 718 | + "### 5.1 获取源码并配置环境\n" | ||
| 719 | + ] | ||
| 720 | + }, | ||
| 721 | + { | ||
| 722 | + "cell_type": "code", | ||
| 723 | + "execution_count": null, | ||
| 724 | + "metadata": {}, | ||
| 725 | + "outputs": [], | ||
| 726 | + "source": [ | ||
| 727 | + "%%bash\n", | ||
| 728 | + "if [ ! -d \"cann-recipes-infer\" ]; then\n", | ||
| 729 | + " git clone https://gitcode.com/cann/cann-recipes-infer.git\n", | ||
| 730 | + "fi\n", | ||
| 731 | + "cd cann-recipes-infer\n", | ||
| 732 | + "\n", | ||
| 733 | + "sed -i 's|cann_path=\"your_cann_pkgs_path\"|cann_path=\"/home/developer/Ascend/cann-9.0.0/aarch64-linux\"|' \\\n", | ||
| 734 | + " executor/scripts/set_env.sh\n", | ||
| 735 | + "\n", | ||
| 736 | + "sed -i '/export ASCEND_HOME_PATH=\\$cann_path/d' \\\n", | ||
| 737 | + " executor/scripts/set_env.sh\n", | ||
| 738 | + "\n", | ||
| 739 | + "pip install -r ./models/qwen/requirements.txt --break-system-packages\n" | ||
| 740 | + ] | ||
| 741 | + }, | ||
| 742 | + { | ||
| 743 | + "cell_type": "markdown", | ||
| 744 | + "metadata": {}, | ||
| 745 | + "source": [ | ||
| 746 | + "### 5.2 复制统一执行器配置模版,写入训练产物路径\n", | ||
| 747 | + "\n", | ||
| 748 | + "复制官方 Qwen3-8B 单卡配置模板,并将 `model_path` 改为本 Notebook 第 3 部分训练产出的 checkpoint 路径:\n" | ||
| 749 | + ] | ||
| 750 | + }, | ||
| 751 | + { | ||
| 752 | + "cell_type": "code", | ||
| 753 | + "execution_count": null, | ||
| 754 | + "metadata": {}, | ||
| 755 | + "outputs": [], | ||
| 756 | + "source": [ | ||
| 757 | + "%%bash\n", | ||
| 758 | + "cd cann-recipes-infer\n", | ||
| 759 | + "cp models/qwen/config/qwen3_8b_1tp.yaml models/qwen/config/qwen3_custom_1tp.yaml\n" | ||
| 760 | + ] | ||
| 761 | + }, | ||
| 762 | + { | ||
| 763 | + "cell_type": "code", | ||
| 764 | + "execution_count": null, | ||
| 765 | + "metadata": {}, | ||
| 766 | + "outputs": [], | ||
| 767 | + "source": [ | ||
| 768 | + "# 将训练产物路径写入运行时 YAML\n", | ||
| 769 | + "import os\n", | ||
| 770 | + "\n", | ||
| 771 | + "yaml_path = \"cann-recipes-infer/models/qwen/config/qwen3_custom_1tp.yaml\"\n", | ||
| 772 | + "checkpoint_path = os.path.abspath(\"./output_qwen3_medical/final\")\n", | ||
| 773 | + "\n", | ||
| 774 | + "with open(yaml_path, \"r\", encoding=\"utf-8\") as f:\n", | ||
| 775 | + " content = f.read()\n", | ||
| 776 | + "\n", | ||
| 777 | + "content = content.replace(\n", | ||
| 778 | + " 'model_path: \"/data/models/origin/Qwen3-8B\"',\n", | ||
| 779 | + " f'model_path: \"{checkpoint_path}\"'\n", | ||
| 780 | + ")\n", | ||
| 781 | + "\n", | ||
| 782 | + "with open(yaml_path, \"w\", encoding=\"utf-8\") as f:\n", | ||
| 783 | + " f.write(content)\n", | ||
| 784 | + "\n", | ||
| 785 | + "print(f\"已将 model_path 写入: {checkpoint_path}\")\n" | ||
| 786 | + ] | ||
| 787 | + }, | ||
| 788 | + { | ||
| 789 | + "cell_type": "markdown", | ||
| 790 | + "metadata": {}, | ||
| 791 | + "source": [ | ||
| 792 | + "### 5.3 补充 lm_head 权重\n", | ||
| 793 | + "\n", | ||
| 794 | + "Qwen 系列模型的 `lm_head` 与 `embed_tokens` 权重共享,训练保存的 checkpoint 中不包含独立的 `lm_head.weight`,接入推理前需要手动补充:\n" | ||
| 795 | + ] | ||
| 796 | + }, | ||
| 797 | + { | ||
| 798 | + "cell_type": "code", | ||
| 799 | + "execution_count": null, | ||
| 800 | + "metadata": {}, | ||
| 801 | + "outputs": [], | ||
| 802 | + "source": [ | ||
| 803 | + "from safetensors.torch import load_file, save_file\n", | ||
| 804 | + "import shutil\n", | ||
| 805 | + "\n", | ||
| 806 | + "CKPT_PATH = os.path.join(checkpoint_path, \"model.safetensors\")\n", | ||
| 807 | + "shutil.copy(CKPT_PATH, CKPT_PATH + \".backup\")\n", | ||
| 808 | + "\n", | ||
| 809 | + "weights = load_file(CKPT_PATH)\n", | ||
| 810 | + "weights[\"lm_head.weight\"] = weights[\"model.embed_tokens.weight\"].clone()\n", | ||
| 811 | + "save_file(weights, CKPT_PATH, metadata={\"format\": \"pt\"})\n", | ||
| 812 | + "\n", | ||
| 813 | + "print(\"已补充 lm_head.weight\")\n" | ||
| 814 | + ] | ||
| 815 | + }, | ||
| 816 | + { | ||
| 817 | + "cell_type": "markdown", | ||
| 818 | + "metadata": {}, | ||
| 819 | + "source": [ | ||
| 820 | + "### 5.4 从 models/qwen 工作目录启动统一执行器\n", | ||
| 821 | + "\n", | ||
| 822 | + "统一执行器运行时会将工作目录切换到对应的模型子目录(如 `models/qwen`)读取配置与资源,在 `cann-recipes-infer` 仓库根目录下执行:\n" | ||
| 823 | + ] | ||
| 824 | + }, | ||
| 825 | + { | ||
| 826 | + "cell_type": "code", | ||
| 827 | + "execution_count": null, | ||
| 828 | + "metadata": {}, | ||
| 829 | + "outputs": [], | ||
| 830 | + "source": [ | ||
| 831 | + "%%bash\n", | ||
| 832 | + "cd cann-recipes-infer\n", | ||
| 833 | + "bash executor/scripts/infer.sh --model qwen --yaml qwen3_custom_1tp.yaml\n" | ||
| 834 | + ] | ||
| 835 | + }, | ||
| 836 | + { | ||
| 837 | + "cell_type": "markdown", | ||
| 838 | + "metadata": {}, | ||
| 839 | + "source": [ | ||
| 840 | + "### 5.5 用自定义领域问题验证效果\n", | ||
| 841 | + "\n", | ||
| 842 | + "默认测试 prompt 与训练领域无关,替换为医学问题后重新执行统一执行器,验证微调效果确实生效:\n" | ||
| 843 | + ] | ||
| 844 | + }, | ||
| 845 | + { | ||
| 846 | + "cell_type": "code", | ||
| 847 | + "execution_count": null, | ||
| 848 | + "metadata": {}, | ||
| 849 | + "outputs": [], | ||
| 850 | + "source": [ | ||
| 851 | + "%%bash\n", | ||
| 852 | + "cd cann-recipes-infer\n", | ||
| 853 | + "cp dataset/default_prompt.json dataset/default_prompt.json.backup\n", | ||
| 854 | + "\n", | ||
| 855 | + "cat > dataset/default_prompt.json << 'EOF'\n", | ||
| 856 | + "{\n", | ||
| 857 | + " \"text\": \"医生,我最近被诊断为糖尿病,应该如何调整饮食?\"\n", | ||
| 858 | + "}\n", | ||
| 859 | + "EOF\n", | ||
| 860 | + "\n", | ||
| 861 | + "bash executor/scripts/infer.sh --model qwen --yaml qwen3_custom_1tp.yaml\n", | ||
| 862 | + "\n", | ||
| 863 | + "mv dataset/default_prompt.json.backup dataset/default_prompt.json\n" | ||
| 864 | + ] | ||
| 865 | + }, | ||
| 866 | + { | ||
| 867 | + "cell_type": "markdown", | ||
| 868 | + "metadata": {}, | ||
| 869 | + "source": [ | ||
| 870 | + "运行结果示例:\n", | ||
| 871 | + "\n", | ||
| 872 | + "**输入**:\n", | ||
| 873 | + "```\n", | ||
| 874 | + "医生,我最近被诊断为糖尿病,应该如何调整饮食?\n", | ||
| 875 | + "```\n", | ||
| 876 | + "\n", | ||
| 877 | + "**输出**:\n", | ||
| 878 | + "```\n", | ||
| 879 | + "<think>\n", | ||
| 880 | + "嗯,用户问的是糖尿病患者如何调整饮食,我需要先回忆一下糖尿病饮食管理的基本原则。首先,我应该从控制总热量开始,因为热量摄入过多会导致血糖升高。然后是碳水化合物的控制,特别是选择低升糖指数的食物,比如全谷物、豆类这些,这样能减缓血糖上升的速度。\n", | ||
| 881 | + "\n", | ||
| 882 | + "接下来是蛋白质和脂肪的摄入,这部分可能需要提到适量,但具体比例可能需要更详细的信息。比如,建议每天摄入多少克蛋白质,脂肪的热量占比是多少。另外,纤维的摄入也很重要,高纤维的食物有助于延缓碳水化合物的吸收,比如蔬菜、水果和全谷物。\n", | ||
| 883 | + "\n", | ||
| 884 | + "然后,用户可能还关心具体的饮食结构,比如三餐的分配,是否有加餐时间,以及如何避免高糖高脂的食物。比如,避免含糖饮料和精制碳水化合物,选择天然甜味剂如蜂蜜或者用天然甜味剂替代。同时,烹饪方式也很关键,比如少油少盐,多用蒸、煮、烤的方式。\n", | ||
| 885 | + "\n", | ||
| 886 | + "另外,用户可能没有提到的点包括饮食记录和监测血糖,这可能也是调整饮食的一部分。不过根据问题,用户主要问的是调整饮食的方法,\n", | ||
| 887 | + "```\n", | ||
| 888 | + "\n", | ||
| 889 | + "> 输出在 `<think>` 推理过程中被截断,未生成最终回答:`qwen3_custom_1tp.yaml` 沿用了官方模板的 `max_new_tokens=256`,256 个 token 的生成上限内还没走出思考阶段。如需完整回答,可在 YAML 里调大 `max_new_tokens`。\n", | ||
| 890 | + "\n", | ||
| 891 | + "**性能实测**(`qwen3_medical_sft.yaml`,1.7B 规模,两次独立复现结果一致):\n", | ||
| 892 | + "\n", | ||
| 893 | + "| 指标 | 实测值 |\n", | ||
| 894 | + "| --- | --- |\n", | ||
| 895 | + "| Prefill | 27.95 ms(复现环境二次实测:27.93 ms) |\n", | ||
| 896 | + "| Decode(平均) | 5.18 ms(复现环境二次实测:5.21 ms,误差在合理范围内) |\n" | ||
| 897 | + ] | ||
| 898 | + }, | ||
| 899 | + { | ||
| 900 | + "cell_type": "markdown", | ||
| 901 | + "metadata": {}, | ||
| 902 | + "source": [ | ||
| 903 | + "至此,完整流程结束:数据准备 → 微调前效果基线 → 构建监督标签 → BF16 全参数 SFT → Transformer 原生推理验证(微调后对比)→ 接入 cann-recipes-infer 统一执行器部署 → 领域问题验证效果。\n" | ||
| 904 | + ] | ||
| 905 | + } | ||
| 906 | + ], | ||
| 907 | + "metadata": { | ||
| 908 | + "kernelspec": { | ||
| 909 | + "display_name": "Python 3", | ||
| 910 | + "language": "python", | ||
| 911 | + "name": "python3" | ||
| 912 | + }, | ||
| 913 | + "language_info": { | ||
| 914 | + "name": "python", | ||
| 915 | + "version": "3.11" | ||
| 916 | + } | ||
| 917 | + }, | ||
| 918 | + "nbformat": 4, | ||
| 919 | + "nbformat_minor": 5 | ||
| 920 | +} | ||
| @@ -0,0 +1,20 @@ | |||
| 1 | +model_config: | ||
| 2 | + model_name: "qwen3_8b" | ||
| 3 | + model_path: "/path/to/your/checkpoint" | ||
| 4 | + exe_mode: "eager" # ["ge_graph", "eager", "npugraph_ex"] | ||
| 5 | + enable_profiler: False # [False, True] | ||
| 6 | + with_ckpt: True # [False, True] | ||
| 7 | + enable_cache_compile: False # [False, True] | ||
| 8 | + enable_static_kernel: False # [False, True] only for npugraph_ex | ||
| 9 | +data_config: | ||
| 10 | + dataset: "default" # ["default", "LongBench"] | ||
| 11 | + input_truncated_len: 512 | ||
| 12 | +parallel_config: | ||
| 13 | + world_size: 1 | ||
| 14 | + attn_tp_size: 1 | ||
| 15 | + moe_tp_size: 1 | ||
| 16 | + embed_tp_size: 1 | ||
| 17 | + lmhead_tp_size: 1 | ||
| 18 | +scheduler_config: | ||
| 19 | + max_new_tokens: 64 | ||
| 20 | + batch_size: 1 | ||
| @@ -0,0 +1,25 @@ | |||
| 1 | +#!/usr/bin/env bash | ||
| 2 | +# 从 cann-recipes-infer 仓库根目录,通过 models/qwen 工作目录启动统一执行器 | ||
| 3 | +# | ||
| 4 | +# 用法: | ||
| 5 | +# bash run_cann_infer.sh <yaml文件名> [cann-recipes-infer仓库路径] | ||
| 6 | +# | ||
| 7 | +# 示例: | ||
| 8 | +# bash run_cann_infer.sh qwen3_custom_1tp.yaml | ||
| 9 | +# bash run_cann_infer.sh qwen3_custom_1tp.yaml /mnt/workspace/gitCode/cann/cann-recipes-infer | ||
| 10 | + | ||
| 11 | +set -euo pipefail | ||
| 12 | + | ||
| 13 | +YAML_FILE="${1:?请提供 YAML 文件名作为第一个参数,如 qwen3_custom_1tp.yaml}" | ||
| 14 | +REPO_PATH="${2:-cann-recipes-infer}" | ||
| 15 | + | ||
| 16 | +if [ ! -d "$REPO_PATH" ]; then | ||
| 17 | + echo "错误:找不到 cann-recipes-infer 仓库路径 $REPO_PATH" | ||
| 18 | + exit 1 | ||
| 19 | +fi | ||
| 20 | + | ||
| 21 | +cd "$REPO_PATH" | ||
| 22 | + | ||
| 23 | +# 统一执行器 executor/scripts/infer.sh 运行时会将工作目录切换到 | ||
| 24 | +# 对应模型子目录(如 models/qwen)读取配置与资源 | ||
| 25 | +bash executor/scripts/infer.sh --model qwen --yaml "$YAML_FILE" | ||
| @@ -0,0 +1,44 @@ | |||
| 1 | +# Validation | ||
| 2 | + | ||
| 3 | +## 1. 环境以及其实测值 | ||
| 4 | + | ||
| 5 | +| 项目 | 实测值 | | ||
| 6 | +| --- | --- | | ||
| 7 | +| 硬件 | 昇腾 Atlas A3(910C),单卡,64GB HBM | | ||
| 8 | +| CANN | 9.0.0 | | ||
| 9 | +| torch_npu | 2.7.1.post4 | | ||
| 10 | +| PyTorch | 2.7.1 | | ||
| 11 | +| transformers | 5.14.1 | | ||
| 12 | +| Python | 3.11 | | ||
| 13 | +| modelscope | 1.22.0(实测环境安装版本,未强制锁定) | | ||
| 14 | +| swanlab | 0.9.0 | | ||
| 15 | + | ||
| 16 | +## 2. 完整 Notebook 实测以及结果数据 | ||
| 17 | + | ||
| 18 | +依据 `qwen3_medical_sft.ipynb` 完整执行,实测数据如下: | ||
| 19 | + | ||
| 20 | +| 阶段 | 指标 | 实测值 | | ||
| 21 | +| --- | --- | --- | | ||
| 22 | +| 数据 | 总样本数 | 2407 条(训练集 2166 / 验证集 241,9:1 切分) | | ||
| 23 | +| 训练 | 训练步数 | ≤680 步(`max_steps=680` 上限,约合 5 个 epoch,配合 `EarlyStoppingCallback` 提前停止,实测收敛于约 650 步附近) | | ||
| 24 | +| 训练 | train/loss(首个 logging step,约) | 1.93 | | ||
| 25 | +| 训练 | train/loss(训练结束时,约) | 1.07 | | ||
| 26 | +| 训练 | eval/loss(训练结束时,约) | 1.16 | | ||
| 27 | +| Transformer 原生推理 | 输出格式 | 正确包含 `<think>...</think>` 思考过程及最终回答 | | ||
| 28 | + | ||
| 29 | +> 上表训练相关数值取自 SwanLab 训练曲线读数(见 `qwen3_medical_sft.ipynb` 中「SwanLab 云端记录」小节的 `eval/loss`、`train/loss` 截图),为近似值。 | ||
| 30 | + | ||
| 31 | +## 3. 实测修正 | ||
| 32 | + | ||
| 33 | +| 问题 | 现象 | 修正方式 | | ||
| 34 | +| --- | --- | --- | | ||
| 35 | +| modelscope 下载数据集接口用错 | `snapshot_download` 下载数据集报 404 | 改用 `git clone` 方式获取数据集 | | ||
| 36 | +| `set_env.sh` 中 `cann_path` 为占位符 | 执行统一执行器报错找不到 `setenv.bash` | 手动替换为真实 CANN 安装路径 | | ||
| 37 | +| checkpoint 缺失 `lm_head.weight` | 接入推理后输出为重复乱码字符 | 手动补充 `lm_head.weight`(复用 `embed_tokens` 权重,见 `qwen3_medical_sft.ipynb` 第 5.3 节) | | ||
| 38 | +| `model_name` 与实际模型规模不一致 | `qwen3_medical_sft.yaml` 中 `model_name` 为 `qwen3_8b`,实际权重为 1.7B | 保留原值,未修改:`cann-recipes-infer` 官方仅为 `qwen3_8b`、`qwen25_7b_instruct` 提供固定的 `model_name` 取值,未见支持自定义规模标签的依据;实测该字段未影响本次加载与推理结果 | | ||
| 39 | +| `trust_remote_code=True` 触发隐藏联网校验,卡死不报错 | 加载 Qwen3-1.7B 模型时(notebook 第 2 部分)进程长时间无响应,`npu-smi info` 显示 HBM 占用长期不变,`top` 显示进程 CPU 占用接近 0%(sleeping 状态),怀疑是联网请求超时挂起,而非真实计算 | 在加载模型代码前加入环境变量强制离线模式,跳过网络校验:`os.environ["HF_HUB_OFFLINE"] = "1"`、`os.environ["TRANSFORMERS_OFFLINE"] = "1"`。该问题是否出现取决于运行环境能否访问 huggingface.co:网络受限(如仅放行 modelscope.cn 等国内域名)的环境下会必现,网络开放的环境下可能感知不到 | | ||
| 40 | +| `cann-recipes-infer/models/qwen/requirements.txt` 锁定 `torch==2.8.0` | 执行 `pip install -r requirements.txt` 后覆盖了预装的 `torch 2.7.1+cpu`,导致 `import torch_npu` 报 `undefined symbol` ABI 不兼容错误 | 执行完依赖安装后需验证 `python3 -c "import torch, torch_npu"` 是否仍能正常导入且版本为 `2.7.1+cpu` / `2.7.1.post4`;若被覆盖,执行 `pip uninstall torch -y` 卸载多装的版本,回退到镜像预装的 `torch 2.7.1+cpu` | | ||
| 41 | +| `torch_npu._C._get_cann_version` 内部读取到非 UTF-8 编码内容 | `source set_env.sh` 设置 `ASCEND_HOME_PATH` 后,`import torch_npu` 触发版本校验逻辑,报 `UnicodeDecodeError: 'utf-8' codec can't decode byte ...`,导致 `RuntimeError: Failed to load the backend extension: torch_npu` | 该函数仅用于版本兼容性提示,非核心推理逻辑必需。对 `torch_npu/npu/utils.py` 中 `get_cann_version` 函数打补丁,用 `try/except UnicodeDecodeError` 包裹调用并返回空字符串兜底(该文件路径通常需要 `sudo` 权限修改) | | ||
| 42 | +| `set_env.sh` 中 `cann_path` 变量身兼两职冲突 | 该变量既用于 `source $cann_path/bin/setenv.bash`(需要带 `aarch64-linux` 层级的路径,因为部分 CANN 安装包中 `cann-9.0.0/bin/setenv.bash` 是指向不存在文件的断链接),又被强制赋值给 `ASCEND_HOME_PATH`(该变量语义上应指向不带 `aarch64-linux` 的上一级目录,`runtime`/`compiler` 等子目录实际存在于该层)。两种用途路径层级要求不一致,导致 `torch_npu` 的 `_cann_package_check` 报 `ASCEND_RUNTIME_PATH` 目录不存在 | 删除 `set_env.sh` 中 `source $cann_path/bin/setenv.bash` 之后那行 `export ASCEND_HOME_PATH=$cann_path` 的强制覆盖,让 `setenv.bash` 脚本内部基于自身路径正确推导出的 `ASCEND_HOME_PATH` 值生效,不再被覆盖 | | ||
| 43 | + | ||
| 44 | +> 以上后三条问题均出现在同一次环境搭建过程中,且具有一定的环境相关性(与所用 CANN 安装包的目录结构、镜像预装的 Python/torch 版本组合有关);其他复现者若使用的 CANN 安装方式与本文实测环境(云端 Ascend Atlas A3 NPU 开发环境,镜像标识 `cann_9.0.0-py3.11-A3-arm`)不同,可能不会触发全部问题,但排查思路可参考本节。 | ||