From b360606763b93cb2ef9fcada442cb3b759b5c243 Mon Sep 17 00:00:00 2001 From: xming521 <1223398803@qq.com> Date: Mon, 21 Apr 2025 21:06:52 +0800 Subject: [PATCH] =?UTF-8?q?=E6=9B=B4=E6=96=B0tests=20=E5=92=8C=20dataset?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- dataset/blocked_words.json | 7 + dataset/res_csv/pt/dataset_info.json | 6 + dataset/res_csv/sft/dataset_info.json | 19 + dataset/test_data.json | 231 ++++++++++ tests/README.md | 39 ++ tests/test_full_pipeline.py | 619 ++++++++++++++++++++++++++ tests/test_weclone_pipeline_mock.py | 306 +++++++++++++ 7 files changed, 1227 insertions(+) create mode 100644 dataset/blocked_words.json create mode 100644 dataset/res_csv/pt/dataset_info.json create mode 100644 dataset/res_csv/sft/dataset_info.json create mode 100644 dataset/test_data.json create mode 100644 tests/README.md create mode 100644 tests/test_full_pipeline.py create mode 100644 tests/test_weclone_pipeline_mock.py diff --git a/dataset/blocked_words.json b/dataset/blocked_words.json new file mode 100644 index 0000000..198c050 --- /dev/null +++ b/dataset/blocked_words.json @@ -0,0 +1,7 @@ +{ + "blocked_words": [ + "例如 姓名", + "例如 地址", + "//....." + ] +} \ No newline at end of file diff --git a/dataset/res_csv/pt/dataset_info.json b/dataset/res_csv/pt/dataset_info.json new file mode 100644 index 0000000..e1ee546 --- /dev/null +++ b/dataset/res_csv/pt/dataset_info.json @@ -0,0 +1,6 @@ +{"wechat-pt":{ + "file_name": "./pt-my.json", + "columns": { + "prompt": "c" + } +}} \ No newline at end of file diff --git a/dataset/res_csv/sft/dataset_info.json b/dataset/res_csv/sft/dataset_info.json new file mode 100644 index 0000000..bb2658e --- /dev/null +++ b/dataset/res_csv/sft/dataset_info.json @@ -0,0 +1,19 @@ +{ + "wechat-sft": { + "file_name": "./sft-my.json", + "columns": { + "prompt": "instruction", + "response": "output", + "system": "system" + } + }, + "wechat-sft-with-history": { + "file_name": "./sft-my.json", + "columns": { + "prompt": "instruction", + "response": "output", + "system": "system", + "history": "history" + } + } +} \ No newline at end of file diff --git a/dataset/test_data.json b/dataset/test_data.json new file mode 100644 index 0000000..057d985 --- /dev/null +++ b/dataset/test_data.json @@ -0,0 +1,231 @@ +{ + "questions": [ + [ + "吃了吗?", + "吃的什么啊", + "好吃吗", + "多少钱啊", + "可以请我吃吗" + ], + [ + "你多大了?" + ], + [ + "你有什么爱好吗?" + ], + [ + "你的理想是什么?", + "你觉得你离你的理想还有多远?" + ], + [ + "你最近在忙什么?", + "工作/学习顺利吗?", + "有什么有趣的事情发生吗?" + ], + [ + "你喜欢看什么类型的电影?", + "最近看过什么好看的电影吗?", + "你最喜欢的电影是什么?" + ], + [ + "你平时喜欢听什么音乐?", + "有推荐的歌手或乐队吗?", + "最近有喜欢的歌曲吗?" + ], + [ + "你喜欢旅游吗?", + "去过哪些地方?", + "最喜欢的旅游地是哪里?" + ], + [ + "你喜欢读书吗?", + "最近在读什么书?", + "最喜欢的书是哪本?" + ], + [ + "你平时喜欢运动吗?", + "喜欢做哪些运动?", + "有固定去锻炼吗?" + ], + [ + "周末一般都做些什么?", + "有没有什么特别的计划?", + "周末喜欢宅在家还是出去玩?" + ], + [ + "你喜欢宠物吗?", + "有养宠物吗?", + "最喜欢什么动物?" + ], + [ + "你喜欢吃什么类型的食物?", + "有推荐的餐厅吗?", + "最喜欢的菜是什么?" + ], + [ + "你喜欢什么样的天气?", + "最喜欢的季节是哪一个?", + "你觉得今天的天气怎么样?" + ], + [ + "你有看电视剧的习惯吗?", + "最近在追哪部剧?", + "最喜欢的电视剧是哪部?" + ], + [ + "你喜欢玩游戏吗?", + "最近在玩什么游戏?", + "有推荐的好玩的游戏吗?" + ], + [ + "你会做饭吗?", + "平时喜欢做哪些菜?", + "有没有特别拿手的菜?" + ], + [ + "你喜欢购物吗?", + "最近买了什么新东西?", + "有推荐的购物网站或店铺吗?" + ], + [ + "你平时怎么放松自己?", + "有特别的解压方式吗?", + "最喜欢的放松活动是什么?" + ], + [ + "你喜欢和朋友出去玩吗?", + "平时会和朋友去哪玩?", + "最近有没有和朋友聚会的计划?" + ], + [ + "你喜欢喝咖啡还是茶?", + "有没有特别喜欢的咖啡馆或茶馆?", + "最喜欢的饮品是什么?" + ], + [ + "你有兄弟姐妹吗?", + "和他们关系怎么样?", + "经常联系吗?" + ], + [ + "你喜欢读什么类型的杂志?", + "最近有看什么有趣的文章吗?", + "有订阅的杂志吗?" + ], + [ + "你喜欢看体育比赛吗?", + "最喜欢的运动项目是什么?", + "有没有特别支持的球队或运动员?" + ], + [ + "你会说其他语言吗?", + "最想学的语言是什么?", + "学习语言有什么技巧吗?" + ], + [ + "你对科技产品感兴趣吗?", + "最近有没有关注什么新科技?", + "最喜欢的电子产品是什么?" + ], + [ + "你喜欢喝什么样的饮料?", + "有没有自己调饮料的习惯?", + "最喜欢的饮品品牌是什么?" + ], + [ + "你平时用社交媒体吗?", + "常用哪些平台?", + "在社交媒体上做什么?" + ], + [ + "你对艺术感兴趣吗?", + "最喜欢的艺术家是谁?", + "有去过哪些艺术展览?" + ], + [ + "你喜欢DIY吗?", + "平时做些什么手工?", + "有没有完成的作品可以分享?" + ], + [ + "你喜欢种植植物吗?", + "有养什么植物?", + "最喜欢的植物是什么?" + ], + [ + "你喜欢拍照吗?", + "喜欢拍什么样的照片?", + "有没有用什么特别的摄影设备?" + ], + [ + "你喜欢听播客吗?", + "常听哪些主题的播客?", + "有没有推荐的播客?" + ], + [ + "你对历史感兴趣吗?", + "最喜欢哪个历史时期?", + "有没有特别喜欢的历史人物?" + ], + [ + "你喜欢画画吗?", + "平时画什么类型的画?", + "有参加过画展吗?" + ], + [ + "你喜欢写作吗?", + "平时写什么类型的文章?", + "有没有发表过作品?" + ], + [ + "你喜欢钓鱼吗?", + "平时去哪里钓鱼?", + "有没有钓到过什么大鱼?" + ], + [ + "你喜欢露营吗?", + "平时会去哪里露营?", + "有没有什么难忘的露营经历?" + ], + [ + "你喜欢摄影吗?", + "最喜欢拍什么题材?", + "有没有特别喜欢的摄影师?" + ], + [ + "你喜欢喝酒吗?", + "喜欢什么类型的酒?", + "有没有推荐的酒吧或品牌?" + ], + [ + "你喜欢滑雪吗?", + "平时去哪里滑雪?", + "有没有什么滑雪技巧分享?" + ], + [ + "你喜欢海边还是山里?", + "最喜欢去哪个地方度假?", + "有没有什么特别推荐的景点?" + ], + [ + "你喜欢参加音乐节吗?", + "参加过哪些音乐节?", + "最喜欢的音乐节是哪一个?" + ], + [ + "你喜欢跑步吗?", + "平时跑多长距离?", + "有没有参加过马拉松?" + ], + [ + "你喜欢参加聚会吗?", + "平时和朋友聚会做什么?", + "有没有什么有趣的聚会游戏?" + ], + [ + "你喜欢收集东西吗?", + "收集什么类型的物品?", + "有没有什么特别的收藏?" + ] + ] +} \ No newline at end of file diff --git a/tests/README.md b/tests/README.md new file mode 100644 index 0000000..d0e6163 --- /dev/null +++ b/tests/README.md @@ -0,0 +1,39 @@ +# WEClone 测试指南 + +本目录包含WEClone项目的测试文件,用于确保项目各个组件正常工作。 + +## 测试文件说明 + +- `test_weclone_pipeline.py`: 全流程测试,按顺序测试数据生成、训练、API服务和模型评估 +- `test_qa_generator.py`: 测试QA生成器功能 + + +## 运行全流程测试 + +要运行完整的测试流程,请执行以下命令: + +```bash +# 在项目根目录下执行 +python -m tests.test_weclone_pipeline +``` + +## 测试流程说明 + +全流程测试按照以下顺序测试项目的主要组件: + +1. **数据生成**:测试 `weclone/data/qa_generator.py` 模块,模拟微信聊天记录的处理和QA对的生成 +2. **模型训练**:测试 `weclone/train/train_sft.py` 模块,模拟使用生成的数据进行模型的SFT训练 +3. **API服务**:测试 `weclone/server/api_service.py` 模块,模拟启动API服务 +4. **模型评估**:测试 `weclone/eval/test_model.py` 模块,模拟对训练后的模型进行评估 + +## 注意事项 + +- 测试使用Python的unittest框架和mock库,模拟各个组件的运行环境和依赖 +- 测试不会修改实际的数据文件或模型文件,所有操作都在临时目录中进行 +- 要运行单独的测试方法,可以使用以下命令: + +```bash +# 例如,只运行QA生成器测试 +python -m unittest tests.test_weclone_pipeline.TestWeclonePipeline.test_qa_generator +``` + diff --git a/tests/test_full_pipeline.py b/tests/test_full_pipeline.py new file mode 100644 index 0000000..84d69d7 --- /dev/null +++ b/tests/test_full_pipeline.py @@ -0,0 +1,619 @@ +import subprocess +import sys +import os +import time +import shutil +import threading # 导入 threading +from typing import Optional, Union, IO # 导入 IO +import torch +from loguru import logger +from subprocess import Popen + +# 配置 Loguru +logger.remove() # 移除默认处理器 +current_time = time.strftime('%Y%m%d_%H%M%S') +log_file_path = os.path.join(os.path.dirname(__file__), f"pipeline_test_{current_time}.log") # 日志文件名包含执行时间 +logger.add(log_file_path, rotation="10 MB", encoding='utf-8', level="DEBUG", enqueue=True) # 文件记录 DEBUG 级别 +logger.add(sys.stdout, colorize=True, format="[test] {time:YYYY-MM-DD HH:mm:ss} | {level.name[0]} | {message}", level="INFO", enqueue=True) # 控制台保持 INFO 级别 + +project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) +logger.info(f"项目根目录: {project_root}") + +qa_script = "weclone/data/qa_generator.py" +train_script = "weclone/train/train_sft.py" +api_service_script = "weclone/server/api_service.py" +eval_script = "weclone/eval/test_model.py" +web_demo_script = "weclone/eval/web_demo.py" + +DEFAULT_TIMEOUT: Optional[Union[int, float]] = 30 +API_STARTUP_WAIT = 20 +API_TERMINATE_WAIT = 15 +WEB_DEMO_STARTUP_WAIT = 20 +WEB_DEMO_TERMINATE_WAIT = 15 + +STEP_QA = "QA 数据生成" +STEP_TRAIN = "SFT 训练" +STEP_COPY_CKPT = "Checkpoint 复制" +STEP_API_START = "API 服务启动" +STEP_EVAL = "模型评估" +STEP_WEB_DEMO = "Web Demo 启动" + +# Mapping from identifiers (script paths or custom keys) to step names +step_identifiers = { + qa_script: STEP_QA, + train_script: STEP_TRAIN, + "copy_checkpoint": STEP_COPY_CKPT, # Custom key for non-script step + api_service_script: STEP_API_START, # Script associated with starting API + eval_script: STEP_EVAL, + web_demo_script: STEP_WEB_DEMO, # Script associated with starting Web Demo +} +# Order for fallback logic +step_order = [STEP_QA, STEP_TRAIN, STEP_COPY_CKPT, STEP_API_START, STEP_EVAL, STEP_WEB_DEMO] + +#todo 需要测试前替换成测试的settings.json 测试完再替换回来 + +class PipelineStepError(Exception): + """自定义异常类,用于表示 Pipeline 步骤执行失败。""" + pass + +# --- 辅助函数:用于在线程中读取和记录流 --- +def log_stream(stream: Optional[IO[str]], log_func): + """读取流并使用指定的 log 函数记录每一行。""" + if stream is None: + return + try: + for line in iter(stream.readline, ''): + if line: + log_func(line.strip()) # 去除末尾换行符 + except ValueError: + # 当 Popen 的 stream 在另一线程中被关闭时,readline 可能会抛出 ValueError + logger.warning("日志流在读取时似乎已被关闭。") + except Exception as e: + # 捕获其他潜在的读取错误 + logger.warning(f"日志流读取时发生未预料的错误: {e}") + finally: + if stream: + try: + stream.close() # 确保流被关闭 + except Exception as close_e: + logger.warning(f"关闭日志流时发生错误: {close_e}") + +# --- 新增:启动日志流线程的辅助函数 --- +def _start_stream_logging_threads(process: Popen, stdout_log_func=logger.info, stderr_log_func=logger.error) -> tuple[threading.Thread, threading.Thread]: + """为给定的进程启动 stdout 和 stderr 的日志记录线程。""" + stdout_thread = threading.Thread( + target=log_stream, + args=(process.stdout, stdout_log_func), + daemon=True + ) + stderr_thread = threading.Thread( + target=log_stream, + args=(process.stderr, stderr_log_func), + daemon=True + ) + stdout_thread.start() + stderr_thread.start() + return stdout_thread, stderr_thread + + +def run_script(script_relative_path: str, timeout: Optional[Union[int, float]] = DEFAULT_TIMEOUT, ignore_timeout_error: bool = False, env: Optional[dict] = None): + """使用 Popen 执行脚本,通过线程实时记录 stdout/stderr 到 loguru。""" + script_full_path = os.path.join(project_root, script_relative_path) + timeout_str = '无限制' if timeout is None else f'{timeout}s' + env_str = f" (环境变量: {env})" if env else "" + logger.info(f"--- 开始执行 (流式): {script_relative_path} (超时: {timeout_str}){env_str} ---") + if not os.path.exists(script_full_path): + error_msg = f"脚本文件不存在 {script_full_path}" + logger.error(error_msg) + raise PipelineStepError(error_msg) + + process: Optional[Popen] = None + stdout_thread: Optional[threading.Thread] = None + stderr_thread: Optional[threading.Thread] = None + + # 准备环境变量 + run_env = os.environ.copy() + if env: + run_env.update(env) + + try: + process = Popen( + [sys.executable, script_full_path], + cwd=project_root, + stdout=subprocess.PIPE, + stderr=subprocess.PIPE, + text=True, + encoding='utf-8', + bufsize=1, # 行缓冲 + env=run_env # 传递环境变量 + ) + + # 使用辅助函数启动日志线程 + stdout_thread, stderr_thread = _start_stream_logging_threads(process, logger.debug, logger.debug) # stdout/stderr 都用 debug + + # 等待子进程完成或超时 + try: + return_code = process.wait(timeout=timeout) + except subprocess.TimeoutExpired: + warn_msg = f"{script_relative_path} 执行超时 ({timeout}s)。" + logger.warning(warn_msg) + # 尝试优雅地关闭流(可能已被 log_stream 关闭) + if process.stdout: process.stdout.close() + if process.stderr: process.stderr.close() + process.kill() # 强制终止超时进程 + logger.warning(f"已强制终止进程 {process.pid}") + # 等待 I/O 线程完成(即使进程被 kill,也要尝试读取剩余输出) + if stdout_thread: stdout_thread.join(timeout=5) + if stderr_thread: stderr_thread.join(timeout=5) + if not ignore_timeout_error: + error_msg = f"{script_relative_path} 执行超时 ({timeout}s) 且未忽略。" + logger.error(error_msg) + raise PipelineStepError(error_msg) + else: + logger.info("--- 根据设置,超时不视为错误,继续执行后续步骤。 ---") + return # 忽略超时,函数正常返回 + + # 等待日志线程完成(确保所有输出都被记录) + if stdout_thread: stdout_thread.join() + if stderr_thread: stderr_thread.join() + + # 检查返回码 + if return_code != 0: + error_msg = f"{script_relative_path} 执行失败,返回码 {return_code}" + logger.error(error_msg) + raise PipelineStepError(error_msg) + else: + logger.success(f"--- {script_relative_path} 执行成功 ---") + + except FileNotFoundError: + error_msg = f"Python 解释器 '{sys.executable}' 或脚本 '{script_full_path}' 未找到。" + logger.error(error_msg) + raise PipelineStepError(error_msg) + except Exception as e: + # 捕获其他潜在错误 (例如 Popen 本身失败) + error_msg = f"执行 {script_relative_path} 时发生意外错误: {e}" + logger.error(error_msg) + # 尝试确保进程和线程被清理 + if process and process.poll() is None: + try: + if process.stdout: process.stdout.close() + if process.stderr: process.stderr.close() + process.kill() + logger.warning(f"因异常 {e},强制终止进程 {process.pid}") + except Exception as kill_e: + logger.error(f"清理过程中强制终止进程失败: {kill_e}") + if stdout_thread and stdout_thread.is_alive(): stdout_thread.join(timeout=1) + if stderr_thread and stderr_thread.is_alive(): stderr_thread.join(timeout=1) + raise PipelineStepError(error_msg) + + +def start_api_service_background() -> Popen: + """在后台启动 API 服务脚本,实时记录启动日志,失败时抛出 PipelineStepError。""" + script_full_path = os.path.join(project_root, api_service_script) + logger.info(f"--- 尝试在后台启动: {api_service_script} ---") + if not os.path.exists(script_full_path): + error_msg = f"脚本文件不存在 {script_full_path}" + logger.error(error_msg) + raise PipelineStepError(error_msg) + + process: Optional[Popen] = None + stdout_thread: Optional[threading.Thread] = None + stderr_thread: Optional[threading.Thread] = None + try: + logger.info(f"启动命令: {[sys.executable, script_full_path]}") + process = Popen( + [sys.executable, script_full_path], + cwd=project_root, + stdout=subprocess.PIPE, + stderr=subprocess.PIPE, + text=True, + encoding='utf-8', + bufsize=1 # 行缓冲 + ) + + # 使用辅助函数启动日志线程 + stdout_thread, stderr_thread = _start_stream_logging_threads(process, logger.debug, logger.debug) # stdout/stderr 都用 debug + + logger.info(f"等待 {API_STARTUP_WAIT} 秒让服务初步启动 (日志将实时显示)...") + time.sleep(API_STARTUP_WAIT) + + # 检查进程是否仍在运行 + if process.poll() is None: + logger.success(f"--- {api_service_script} 似乎已在后台启动 (进程 PID: {process.pid}) ---") + # 注意:不 join 日志线程,让它们继续运行 + return process + else: + # 进程过早退出 + logger.error(f"{api_service_script} 启动后在 {API_STARTUP_WAIT} 秒内过早退出,返回码 {process.returncode}") + # 尝试等待日志线程结束以捕获最后输出 + if stdout_thread: stdout_thread.join(timeout=2) + if stderr_thread: stderr_thread.join(timeout=2) + # 读取 communicate 获取可能遗漏的最终输出 (虽然理论上线程应该读完了) + try: + # 设置短超时,因为进程已退出,communicate 应该立即返回 + stdout, stderr = process.communicate(timeout=1) + except subprocess.TimeoutExpired: + logger.warning("等待 communicate 超时,可能没有更多输出了。") + stdout, stderr = "", "" # 假设没有更多输出 + except Exception as comm_e: + logger.warning(f"调用 communicate 获取最后输出时出错: {comm_e}") + stdout, stderr = "", "" + + error_message = f'''--- EARLY EXIT STDOUT --- + {stdout} + --- EARLY EXIT STDERR --- + {stderr}''' + logger.error(error_message) + raise PipelineStepError(f"{api_service_script} 启动失败并过早退出。") + + except FileNotFoundError: + error_msg = f"Python 解释器 '{sys.executable}' 或脚本 '{script_full_path}' 未找到。" + logger.error(error_msg) + raise PipelineStepError(error_msg) + except Exception as e: + # 捕获其他启动错误 + error_msg = f"启动 {api_service_script} 时发生意外错误: {e}" + logger.error(error_msg) + if process and process.poll() is None: + logger.warning("捕获到异常,尝试强制终止进程...") + try: + if process.stdout: process.stdout.close() + if process.stderr: process.stderr.close() + process.kill() + except Exception as kill_e: logger.error(f"强制终止进程时出错: {kill_e}") + # 尝试join线程 + if stdout_thread and stdout_thread.is_alive(): stdout_thread.join(timeout=1) + if stderr_thread and stderr_thread.is_alive(): stderr_thread.join(timeout=1) + raise PipelineStepError(error_msg) + +def stop_api_service(process: Optional[Popen]): + """停止指定的 API 服务进程。""" + if process and process.poll() is None: + logger.info(f"--- 尝试停止 API 服务 (PID: {process.pid}) ---") + try: + # 先关闭流,再终止进程,避免 log_stream 线程因流关闭而出错 + if process.stdout: process.stdout.close() + if process.stderr: process.stderr.close() + process.terminate() + logger.info(f"发送终止信号,等待 {API_TERMINATE_WAIT} 秒让服务优雅终止...") + try: + process.wait(timeout=API_TERMINATE_WAIT) # 等待进程实际结束 + logger.info(f"API 服务进程已优雅终止,返回码: {process.returncode}") + except subprocess.TimeoutExpired: + logger.warning(f"优雅终止超时 ({API_TERMINATE_WAIT}s),强制终止进程...") + process.kill() + process.wait() # 等待强制终止完成 + logger.info("API 服务进程已被强制终止。") + # communicate() 在这里可能不再需要,因为我们主动关闭了流并且等待了进程 + # 如果需要最后的输出,可能需要在 kill/terminate 前读取 + except Exception as e: + logger.error(f"停止 API 服务时发生错误: {e}") + # 如果停止过程中出错,尝试强制kill + if process.poll() is None: + logger.warning("停止过程中出现错误,尝试强制终止...") + try: + process.kill() + process.wait() + except Exception as kill_e: + logger.error(f"停止过程中强制终止进程时出错: {kill_e}") + + elif process: + logger.info(f"--- API 服务进程 (PID: {process.pid}) 在尝试停止前已经退出。 ---") + else: + logger.debug("--- 无需停止 API 服务 (进程不存在或已为 None) ---") + +def start_web_demo_background() -> Popen: + """在后台启动 Web Demo 脚本,实时记录启动日志,失败时抛出 PipelineStepError。""" + script_full_path = os.path.join(project_root, web_demo_script) + logger.info(f"--- 尝试在后台启动: {web_demo_script} ---") + if not os.path.exists(script_full_path): + error_msg = f"脚本文件不存在 {script_full_path}" + logger.error(error_msg) + raise PipelineStepError(error_msg) + + process: Optional[Popen] = None + stdout_thread: Optional[threading.Thread] = None + stderr_thread: Optional[threading.Thread] = None + try: + logger.info(f"启动命令: {[sys.executable, script_full_path]}") + process = Popen( + [sys.executable, script_full_path], + cwd=project_root, + stdout=subprocess.PIPE, + stderr=subprocess.PIPE, + text=True, + encoding='utf-8', + bufsize=1 # 行缓冲 + ) + + # 使用辅助函数启动日志线程 (stdout/stderr 都用 info) + stdout_thread, stderr_thread = _start_stream_logging_threads(process, logger.debug, logger.debug) # stdout/stderr 都用 debug + + + logger.info(f"等待 {WEB_DEMO_STARTUP_WAIT} 秒让 Web Demo 初步启动 (日志将实时显示)...") + time.sleep(WEB_DEMO_STARTUP_WAIT) + + # 检查进程是否仍在运行 + if process.poll() is None: + logger.success(f"--- {web_demo_script} 似乎已在后台启动 (进程 PID: {process.pid}) ---") + # 注意:不 join 日志线程 + return process + else: + # 进程过早退出 + logger.error(f"{web_demo_script} 启动后在 {WEB_DEMO_STARTUP_WAIT} 秒内过早退出,返回码 {process.returncode}") + if stdout_thread: stdout_thread.join(timeout=2) + if stderr_thread: stderr_thread.join(timeout=2) + try: + stdout, stderr = process.communicate(timeout=1) + except subprocess.TimeoutExpired: + logger.warning("等待 communicate 超时,可能没有更多输出了。") + stdout, stderr = "", "" + except Exception as comm_e: + logger.warning(f"调用 communicate 获取最后输出时出错: {comm_e}") + stdout, stderr = "", "" + error_message = f'''--- EARLY EXIT STDOUT --- + {stdout} + --- EARLY EXIT STDERR --- + {stderr}''' + logger.error(error_message) + raise PipelineStepError(f"{web_demo_script} 启动失败并过早退出。") + + except FileNotFoundError: + error_msg = f"Python 解释器 '{sys.executable}' 或脚本 '{script_full_path}' 未找到。" + logger.error(error_msg) + raise PipelineStepError(error_msg) + except Exception as e: + error_msg = f"启动 {web_demo_script} 时发生意外错误: {e}" + logger.error(error_msg) + if process and process.poll() is None: + logger.warning("捕获到异常,尝试强制终止进程...") + try: + if process.stdout: process.stdout.close() + if process.stderr: process.stderr.close() + process.kill() + except Exception as kill_e: logger.error(f"强制终止进程时出错: {kill_e}") + if stdout_thread and stdout_thread.is_alive(): stdout_thread.join(timeout=1) + if stderr_thread and stderr_thread.is_alive(): stderr_thread.join(timeout=1) + raise PipelineStepError(error_msg) + +def stop_web_demo(process: Optional[Popen]): + """停止指定的 Web Demo 进程。""" + if process and process.poll() is None: + logger.info(f"--- 尝试停止 Web Demo 服务 (PID: {process.pid}) ---") + try: + # 先关闭流,再终止进程 + if process.stdout: process.stdout.close() + if process.stderr: process.stderr.close() + process.terminate() + logger.info(f"发送终止信号,等待 {WEB_DEMO_TERMINATE_WAIT} 秒让服务优雅终止...") + try: + process.wait(timeout=WEB_DEMO_TERMINATE_WAIT) + logger.info(f"Web Demo 服务进程已优雅终止,返回码: {process.returncode}") + except subprocess.TimeoutExpired: + logger.warning(f"优雅终止超时 ({WEB_DEMO_TERMINATE_WAIT}s),强制终止进程...") + process.kill() + process.wait() + logger.info("Web Demo 服务进程已被强制终止。") + except Exception as e: + logger.error(f"停止 Web Demo 服务时发生错误: {e}") + if process.poll() is None: + logger.warning("停止过程中出现错误,尝试强制终止...") + try: + process.kill() + process.wait() + except Exception as kill_e: + logger.error(f"停止过程中强制终止进程时出错: {kill_e}") + elif process: + logger.info(f"--- Web Demo 服务进程 (PID: {process.pid}) 在尝试停止前已经退出。 ---") + else: + logger.debug("--- 无需停止 Web Demo 服务 (进程不存在或已为 None) ---") + + +if __name__ == "__main__": + logger.info("="*20 + " 开始执行 WeClone Pipeline 脚本 " + "="*20) + + is_cuda_available = torch.cuda.is_available() + logger.info("--- CUDA 可用性检查 ---") + if is_cuda_available: + gpu_count = torch.cuda.device_count() + logger.success(f"CUDA 可用 (找到 {gpu_count} 个 GPU)") + for i in range(gpu_count): + logger.info(f" - GPU {i}: {torch.cuda.get_device_name(i)}") + else: + logger.warning("CUDA 不可用,将使用 CPU (如果适用)。") + logger.info("-" * 25) + + steps_completed = [] + api_process: Optional[Popen] = None + web_demo_process: Optional[Popen] = None + + # 设置哪些步骤需要运行 + run_qa = True + run_train = True + run_copy_checkpoint = True # 依赖于 run_train + run_api = True + run_eval = True # 依赖于 run_api + run_web_demo = True # 不依赖于 run_api + + try: + # 步骤 1: QA Generator + if run_qa: + logger.info("-" * 10 + " 步骤 1: QA 数据生成 " + "-" * 10) + run_script(qa_script) + steps_completed.append(f"{STEP_QA}: 成功") + else: + logger.info(f"{STEP_QA}: 跳过 (配置)") + steps_completed.append(f"{STEP_QA}: 跳过") + + # 步骤 2: Train SFT + if run_train: + logger.info("-" * 10 + " 步骤 2: SFT 训练 " + "-" * 10) + # 删除 model_output 目录 + model_output_dir = os.path.join(project_root, "model_output") + if os.path.exists(model_output_dir): + logger.info(f"删除现有的 model_output 目录: {model_output_dir}") + try: + shutil.rmtree(model_output_dir) + logger.success("成功删除 model_output 目录") + except Exception as e: + logger.error(f"删除 model_output 目录时出错: {e}") + + # 尝试禁用 tqdm + run_script(train_script, timeout=DEFAULT_TIMEOUT, ignore_timeout_error=True, env={'TQDM_DISABLE': '1'}) + steps_completed.append(f"{STEP_TRAIN}: 成功或超时跳过") + + # 步骤 2.1: 复制 Checkpoint (只有在训练运行后才可能执行) + if run_copy_checkpoint: + logger.info("-" * 10 + " 步骤 2.1: 复制 Checkpoint 到 model_output " + "-" * 10) + source_dir = os.path.join(project_root, "model_output", "checkpoint-2") + dest_dir = os.path.join(project_root, "model_output") + if os.path.isdir(source_dir): + try: + logger.info(f"开始将 {source_dir} 的内容复制到 {dest_dir}...") + shutil.copytree(source_dir, dest_dir, dirs_exist_ok=True) + logger.success(f"--- {STEP_COPY_CKPT} 成功 ---") + steps_completed.append(f"{STEP_COPY_CKPT}: 成功") + except Exception as e: + # Embed identifier in the error for the except block + error_msg = f"{STEP_COPY_CKPT} 时发生错误: {e}" + logger.error(error_msg) + # Add a unique marker to identify this step in the except block + raise PipelineStepError(f"{error_msg} ###step_id:copy_checkpoint###") + else: + logger.warning(f"源 Checkpoint 目录 {source_dir} 不存在或不是目录,跳过复制。") + steps_completed.append(f"{STEP_COPY_CKPT}: 跳过 (源不存在)") + # raise PipelineStepError(f"必需的源 Checkpoint 目录 {source_dir} 不存在") # 如果必须,取消此行注释 + else: + logger.info(f"{STEP_COPY_CKPT}: 跳过 (配置)") + steps_completed.append(f"{STEP_COPY_CKPT}: 跳过 (配置)") + + else: + logger.info(f"{STEP_TRAIN}: 跳过 (配置)") + steps_completed.append(f"{STEP_TRAIN}: 跳过 (配置)") + logger.info(f"{STEP_COPY_CKPT}: 跳过 (训练未运行)") + steps_completed.append(f"{STEP_COPY_CKPT}: 跳过 (训练未运行)") + + + # 步骤 3: Start API Service + if run_api: + logger.info("-" * 10 + " 步骤 3: 启动 API 服务 " + "-" * 10) + api_process = start_api_service_background() + steps_completed.append(f"{STEP_API_START}: 成功") + else: + logger.info(f"{STEP_API_START}: 跳过 (配置)") + steps_completed.append(f"{STEP_API_START}: 跳过 (配置)") + + + # 步骤 4: Eval Model (依赖 API 服务) + if run_eval: + if not run_api: + logger.info("-" * 10 + " 步骤 4: 模型评估 " + "-" * 10) + logger.warning("--- 因 API 服务配置为不运行,跳过执行: weclone/eval/test_model.py ---") + steps_completed.append(f"{STEP_EVAL}: 跳过 (API未配置运行)") + elif api_process is None or api_process.poll() is not None: # 检查进程是否已退出 + error_msg = "尝试运行评估,但 API 服务进程不存在或已退出。" + logger.error(error_msg) + raise PipelineStepError(error_msg) + else: + logger.info("-" * 10 + " 步骤 4: 模型评估 " + "-" * 10) + # 在调用评估脚本时禁用 tqdm + run_script(eval_script, timeout=9999, env={'TQDM_DISABLE': '1'}) + steps_completed.append(f"{STEP_EVAL}: 成功") + stop_api_service(api_process) # 评估完成后停止API + api_process = None # 标记为已停止 + else: + logger.info(f"{STEP_EVAL}: 跳过 (配置)") + steps_completed.append(f"{STEP_EVAL}: 跳过 (配置)") + if api_process: # 如果API在运行但评估被跳过,也停止API + logger.info("评估被跳过,停止 API 服务...") + stop_api_service(api_process) + api_process = None + + + # 步骤 5: Start Web Demo (不依赖 API 服务) + if run_web_demo: + logger.info("-" * 10 + " 步骤 5: 启动 Web Demo " + "-" * 10) + web_demo_process = start_web_demo_background() + steps_completed.append(f"{STEP_WEB_DEMO}: 成功") + logger.info("--- Web Demo 已启动,测试流程继续... ---") + else: + logger.info(f"{STEP_WEB_DEMO}: 跳过 (配置)") + steps_completed.append(f"{STEP_WEB_DEMO}: 跳过 (配置)") + + # Pipeline 成功完成所有请求的步骤 + logger.info("="*20 + " Pipeline 执行摘要 " + "="*20) + for step in steps_completed: + logger.info(f"- {step}") + logger.success("✅ 所有请求执行的 Pipeline 步骤均成功完成!") + + skipped_steps = [s for s in steps_completed if "跳过" in s] + if skipped_steps: + logger.warning("注意: 以下步骤被设置为跳过,如需执行请修改脚本顶部的 run_xxx 变量:") + for skipped in skipped_steps: + logger.warning(f" - {skipped.split(':')[0]}") + + + except PipelineStepError as e: + logger.error("="*20 + " Pipeline 执行失败 " + "="*20) + + failing_step = "未知步骤" + error_details = str(e) + cleaned_error_details = error_details # Store original/cleaned details for logging + + # Attempt 1: Check for explicit marker (e.g., from copy_checkpoint) + marker_prefix = "###step_id:" + marker_suffix = "###" + marker_start = error_details.find(marker_prefix) + if marker_start != -1: + marker_end = error_details.find(marker_suffix, marker_start + len(marker_prefix)) + if marker_end != -1: + step_id = error_details[marker_start + len(marker_prefix):marker_end] + failing_step = step_identifiers.get(step_id, f"未知标记 ({step_id})") + # Clean the marker from the displayed error details + cleaned_error_details = error_details[:marker_start].strip() + + # Attempt 2: Check for known script paths in the error message if marker not found + if failing_step == "未知步骤": + found_script = False + # Iterate through potential script paths stored as keys in step_identifiers + for identifier, step_name in step_identifiers.items(): + # Check if the identifier looks like a path and is in the error message + if isinstance(identifier, str) and ('/' in identifier or '\\\\' in identifier) and identifier in error_details: + failing_step = step_name + found_script = True + break # Found the most likely script + + # Attempt 3: Fallback based on last completed step (if still unknown) + if failing_step == "未知步骤": + if steps_completed: + last_completed = steps_completed[-1].split(':')[0] + try: + last_completed_index = step_order.index(last_completed) + if last_completed_index + 1 < len(step_order): + # Assume the next step in the defined order failed + failing_step = step_order[last_completed_index + 1] + " (推断)" + else: + failing_step = "Pipeline末尾或未知 (推断)" # Error after the last known step + except ValueError: + # Last completed step name wasn't found in our defined order + failing_step = f"未知 (最后完成: {last_completed})" + else: + failing_step = "初始化期间" # No steps completed + + logger.error(f"错误发生在步骤: {failing_step}") + logger.error(f"错误详情: {cleaned_error_details}") # Log the cleaned error details + logger.info("--- 已完成步骤 ---") + for step in steps_completed: + logger.info(f"- {step}") + logger.error("="*50) + sys.exit(1) # 测试失败时退出码为 1 + + finally: + logger.info("--- Pipeline 结束,开始清理后台服务 ---") + # 确保在 finally 块中总是尝试停止服务 + stop_web_demo(web_demo_process) + stop_api_service(api_process) # 即使评估步骤停止了它,这里也尝试停止,无害 + logger.info("--- 后台服务清理完成 ---") + + # 如果 Pipeline 成功,确保退出码为 0 + sys.exit(0) \ No newline at end of file diff --git a/tests/test_weclone_pipeline_mock.py b/tests/test_weclone_pipeline_mock.py new file mode 100644 index 0000000..b30bbc1 --- /dev/null +++ b/tests/test_weclone_pipeline_mock.py @@ -0,0 +1,306 @@ +import os +import sys +import json +import shutil +import unittest +import tempfile +from unittest.mock import patch, MagicMock +import pandas as pd + +# 添加项目根目录到系统路径 +sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) + +# 导入需要测试的模块 +from weclone.data.qa_generator import DataProcessor +from weclone.utils.config import load_config + + +class TestWeclonePipeline(unittest.TestCase): + @classmethod + def setUpClass(cls): + """设置测试环境""" + # 创建临时目录用于测试 + cls.test_dir = tempfile.mkdtemp() + cls.test_data_dir = os.path.join(cls.test_dir, "data") + cls.test_model_dir = os.path.join(cls.test_dir, "model_output") + cls.test_eval_dir = os.path.join(cls.test_dir, "eval_output") + + # 创建必要的目录 + os.makedirs(cls.test_data_dir, exist_ok=True) + os.makedirs(cls.test_model_dir, exist_ok=True) + os.makedirs(cls.test_eval_dir, exist_ok=True) + + # 创建测试数据集结构 + cls.csv_folder = os.path.join(cls.test_data_dir, "csv") + os.makedirs(cls.csv_folder, exist_ok=True) + + # 创建示例聊天文件夹和CSV文件 + chat_folder = os.path.join(cls.csv_folder, "test_chat") + os.makedirs(chat_folder, exist_ok=True) + + # 创建简单的测试CSV数据 + cls._create_test_csv(os.path.join(chat_folder, "test_chat.csv")) + + # 创建测试用的settings.json + cls._create_test_settings() + + # 创建测试用的test_data.json用于模型评估 + cls._create_test_eval_data() + + @classmethod + def tearDownClass(cls): + """清理测试环境""" + # 删除临时目录 + shutil.rmtree(cls.test_dir, ignore_errors=True) + + @classmethod + def _create_test_csv(cls, file_path): + """创建测试用CSV文件""" + import pandas as pd + + # 创建简单的聊天记录数据 + data = { + "id": list(range(1, 5)), + "MsgSvrID": list(range(1001, 1005)), + "type": ["1", "1", "1", "1"], # 文本类型 + "is_sender": [0, 1, 0, 1], # 0=对方发送,1=自己发送 + "talker": ["test_user", "me", "test_user", "me"], + "room_name": ["", "", "", ""], + "content": ["你好,请问你是谁?", "我是你的微信助手", "你能帮我做什么?", "我可以回答问题,提供信息和帮助你完成各种任务"], + "src": ["", "", "", ""], + "CreateTime": [1609459200, 1609459220, 1609459240, 1609459260] # 时间戳 + } + + # 创建DataFrame并保存为CSV + df = pd.DataFrame(data) + df.to_csv(file_path, index=False) + + @classmethod + def _create_test_settings(cls): + """创建测试用的settings.json""" + # 简化版的设置文件,只包含测试所需的最小配置 + settings = { + "train_sft_args": { + "stage": "sft", + "dataset": "wechat-sft", + "dataset_dir": cls.test_data_dir + "/res_csv/sft", + "lora_target": "query_key_value", + "lora_rank": 4, + "lora_dropout": 0.5, + "overwrite_cache": True, + "per_device_train_batch_size": 1, + "gradient_accumulation_steps": 1, + "lr_scheduler_type": "cosine", + "logging_steps": 1, + "save_steps": 1, + "learning_rate": 0.0001, + "num_train_epochs": 1, + "plot_loss": False, + "fp16": False + }, + "infer_args": { + "repetition_penalty": 1.2, + "temperature": 0.5, + "max_length": 50, + "top_p": 0.65 + }, + "make_dataset_args": { + "single_combine_strategy": "time_window", + "qa_match_strategy": "time_window", + "single_combine_time_window": 2, + "qa_match_time_window": 5, + "prompt_with_history": False + }, + "common_args": { + "model_name_or_path": "./chatglm3-6b", # 假设已有模型 + "adapter_name_or_path": cls.test_model_dir, + "template": "chatglm3-weclone", + "finetuning_type": "lora", + "trust_remote_code": True + } + } + + # 保存到临时目录 + with open(os.path.join(cls.test_dir, "settings.json"), "w", encoding="utf-8") as f: + json.dump(settings, f, indent=4) + + @classmethod + def _create_test_eval_data(cls): + """创建测试用的评估数据""" + test_data = { + "questions": [ + ["你好", "你是谁"], + ["你能做什么"] + ] + } + + # 确保目录存在 + data_dir = os.path.join(cls.test_dir, "data") + os.makedirs(data_dir, exist_ok=True) + + # 保存测试数据 + with open(os.path.join(data_dir, "test_data.json"), "w", encoding="utf-8") as f: + json.dump(test_data, f, ensure_ascii=False, indent=4) + + @patch('weclone.data.qa_generator.DataProcessor.get_csv_files') + @patch('weclone.data.qa_generator.DataProcessor.load_csv') + @patch('weclone.data.qa_generator.DataProcessor.save_result') + def test_qa_generator(self, mock_save_result, mock_load_csv, mock_get_csv_files): + """测试QA生成器""" + print("\n测试QA生成器...") + + # 准备模拟数据 + from weclone.data.models import ChatMessage + mock_get_csv_files.return_value = ["test_csv_file.csv"] + + # 模拟从CSV加载的消息 + mock_messages = [ + ChatMessage(id=1, MsgSvrID=1001, type_name="文本", is_sender=0, + talker="test_user", room_name="", msg="你好,请问你是谁?", + src="", CreateTime=pd.Timestamp(1609459200, unit='s')), + ChatMessage(id=2, MsgSvrID=1002, type_name="文本", is_sender=1, + talker="me", room_name="", msg="我是你的微信助手", + src="", CreateTime=pd.Timestamp(1609459220, unit='s')), + ChatMessage(id=3, MsgSvrID=1003, type_name="文本", is_sender=0, + talker="test_user", room_name="", msg="你能帮我做什么?", + src="", CreateTime=pd.Timestamp(1609459240, unit='s')), + ChatMessage(id=4, MsgSvrID=1004, type_name="文本", is_sender=1, + talker="me", room_name="", msg="我可以回答问题,提供信息和帮助你完成各种任务", + src="", CreateTime=pd.Timestamp(1609459260, unit='s')) + ] + mock_load_csv.return_value = mock_messages + + # 创建DataProcessor实例 + with patch('weclone.utils.config.load_config') as mock_load_config: + # 模拟配置 + mock_config = { + "single_combine_strategy": "time_window", + "qa_match_strategy": "time_window", + "single_combine_time_window": 2, + "qa_match_time_window": 5, + "prompt_with_history": False + } + mock_load_config.return_value = mock_config + + # 执行QA生成 + processor = DataProcessor() + processor.csv_folder = self.csv_folder # 设置为测试目录 + processor.main() + + # 验证是否调用了预期的方法 + mock_get_csv_files.assert_called_once() + mock_load_csv.assert_called_once() + mock_save_result.assert_called_once() + + # 验证结果格式 + # 获取保存的结果 + call_args = mock_save_result.call_args[0][0] + self.assertTrue(isinstance(call_args, list)) + self.assertEqual(len(call_args), 2) # 应该有两个QA对 + + # 验证QA对的结构 + for qa in call_args: + self.assertTrue("instruction" in qa) + self.assertTrue("output" in qa) + + print("QA生成器测试成功") + + def test_train_sft(self): + """测试SFT训练过程""" + print("\n测试SFT训练过程...") + # 由于训练需要实际的模型和数据,这里我们只模拟调用 + + with patch('llamafactory.train.tuner.run_exp') as mock_run_exp: + # 导入训练模块并运行 + from weclone.train.train_sft import run_exp + + + # 验证是否正确调用了训练函数 + self.assertTrue(mock_run_exp.called) + print("SFT训练过程测试成功") + + def test_api_service(self): + """测试API服务""" + print("\n测试API服务...") + + # 模拟服务器进程 + with patch('uvicorn.run') as mock_run: + # 导入API服务模块 + from weclone.server.api_service import main, create_app, ChatModel + + # 模拟配置和模型 + with patch('weclone.utils.config.load_config') as mock_load_config: + mock_config = {"model_path": "test_model_path"} + mock_load_config.return_value = mock_config + + # 模拟ChatModel + with patch('llamafactory.chat.ChatModel') as MockChatModel: + mock_chat_model = MagicMock() + MockChatModel.return_value = mock_chat_model + + # 运行API服务 + main() + + # 验证服务是否正确启动 + mock_run.assert_called_once() + call_args = mock_run.call_args[1] + self.assertEqual(call_args["host"], "0.0.0.0") + self.assertEqual(call_args["port"], 8005) # 默认端口 + self.assertEqual(call_args["workers"], 1) + + print("API服务测试成功") + + def test_model_evaluation(self): + """测试模型评估""" + print("\n测试模型评估...") + + # 模拟OpenAI API调用 + with patch('openai.ChatCompletion.create') as mock_create: + # 设置模拟返回值 + mock_response = MagicMock() + mock_response.choices = [MagicMock()] + mock_response.choices[0].message.content = "这是模型的测试回复" + mock_create.return_value = mock_response + + # 运行评估脚本 + with patch('builtins.open', create=True) as mock_open: + # 模拟打开测试数据文件 + test_data_content = '{"questions": [["你好", "你是谁"], ["你能做什么"]]}' + mock_file = MagicMock() + mock_file.read.return_value = test_data_content + mock_open.return_value.__enter__.return_value = mock_file + + # 导入并运行评估模块 + from weclone.eval.test_model import main + + # 执行评估 + main() + + # 验证API调用次数(应该是测试问题的数量) + self.assertEqual(mock_create.call_count, 3) # 3个测试问题 + + print("模型评估测试成功") + + def test_full_pipeline(self): + """测试完整流程""" + print("\n测试完整流程...") + + # 这个测试方法会依次调用上面的各个测试方法,模拟完整的流程 + + # 1. 测试QA生成器 + self.test_qa_generator() + + # 2. 测试SFT训练 + self.test_train_sft() + + # 3. 测试API服务 + self.test_api_service() + + # 4. 测试模型评估 + self.test_model_evaluation() + + print("完整流程测试完成") + + +if __name__ == "__main__": + unittest.main() \ No newline at end of file