mirror of
https://github.com/ooyinet/WeClone.git
synced 2026-08-28 22:16:52 +08:00
更新tests 和 dataset
This commit is contained in:
@@ -0,0 +1,7 @@
|
||||
{
|
||||
"blocked_words": [
|
||||
"例如 姓名",
|
||||
"例如 地址",
|
||||
"//....."
|
||||
]
|
||||
}
|
||||
@@ -0,0 +1,6 @@
|
||||
{"wechat-pt":{
|
||||
"file_name": "./pt-my.json",
|
||||
"columns": {
|
||||
"prompt": "c"
|
||||
}
|
||||
}}
|
||||
@@ -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"
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,231 @@
|
||||
{
|
||||
"questions": [
|
||||
[
|
||||
"吃了吗?",
|
||||
"吃的什么啊",
|
||||
"好吃吗",
|
||||
"多少钱啊",
|
||||
"可以请我吃吗"
|
||||
],
|
||||
[
|
||||
"你多大了?"
|
||||
],
|
||||
[
|
||||
"你有什么爱好吗?"
|
||||
],
|
||||
[
|
||||
"你的理想是什么?",
|
||||
"你觉得你离你的理想还有多远?"
|
||||
],
|
||||
[
|
||||
"你最近在忙什么?",
|
||||
"工作/学习顺利吗?",
|
||||
"有什么有趣的事情发生吗?"
|
||||
],
|
||||
[
|
||||
"你喜欢看什么类型的电影?",
|
||||
"最近看过什么好看的电影吗?",
|
||||
"你最喜欢的电影是什么?"
|
||||
],
|
||||
[
|
||||
"你平时喜欢听什么音乐?",
|
||||
"有推荐的歌手或乐队吗?",
|
||||
"最近有喜欢的歌曲吗?"
|
||||
],
|
||||
[
|
||||
"你喜欢旅游吗?",
|
||||
"去过哪些地方?",
|
||||
"最喜欢的旅游地是哪里?"
|
||||
],
|
||||
[
|
||||
"你喜欢读书吗?",
|
||||
"最近在读什么书?",
|
||||
"最喜欢的书是哪本?"
|
||||
],
|
||||
[
|
||||
"你平时喜欢运动吗?",
|
||||
"喜欢做哪些运动?",
|
||||
"有固定去锻炼吗?"
|
||||
],
|
||||
[
|
||||
"周末一般都做些什么?",
|
||||
"有没有什么特别的计划?",
|
||||
"周末喜欢宅在家还是出去玩?"
|
||||
],
|
||||
[
|
||||
"你喜欢宠物吗?",
|
||||
"有养宠物吗?",
|
||||
"最喜欢什么动物?"
|
||||
],
|
||||
[
|
||||
"你喜欢吃什么类型的食物?",
|
||||
"有推荐的餐厅吗?",
|
||||
"最喜欢的菜是什么?"
|
||||
],
|
||||
[
|
||||
"你喜欢什么样的天气?",
|
||||
"最喜欢的季节是哪一个?",
|
||||
"你觉得今天的天气怎么样?"
|
||||
],
|
||||
[
|
||||
"你有看电视剧的习惯吗?",
|
||||
"最近在追哪部剧?",
|
||||
"最喜欢的电视剧是哪部?"
|
||||
],
|
||||
[
|
||||
"你喜欢玩游戏吗?",
|
||||
"最近在玩什么游戏?",
|
||||
"有推荐的好玩的游戏吗?"
|
||||
],
|
||||
[
|
||||
"你会做饭吗?",
|
||||
"平时喜欢做哪些菜?",
|
||||
"有没有特别拿手的菜?"
|
||||
],
|
||||
[
|
||||
"你喜欢购物吗?",
|
||||
"最近买了什么新东西?",
|
||||
"有推荐的购物网站或店铺吗?"
|
||||
],
|
||||
[
|
||||
"你平时怎么放松自己?",
|
||||
"有特别的解压方式吗?",
|
||||
"最喜欢的放松活动是什么?"
|
||||
],
|
||||
[
|
||||
"你喜欢和朋友出去玩吗?",
|
||||
"平时会和朋友去哪玩?",
|
||||
"最近有没有和朋友聚会的计划?"
|
||||
],
|
||||
[
|
||||
"你喜欢喝咖啡还是茶?",
|
||||
"有没有特别喜欢的咖啡馆或茶馆?",
|
||||
"最喜欢的饮品是什么?"
|
||||
],
|
||||
[
|
||||
"你有兄弟姐妹吗?",
|
||||
"和他们关系怎么样?",
|
||||
"经常联系吗?"
|
||||
],
|
||||
[
|
||||
"你喜欢读什么类型的杂志?",
|
||||
"最近有看什么有趣的文章吗?",
|
||||
"有订阅的杂志吗?"
|
||||
],
|
||||
[
|
||||
"你喜欢看体育比赛吗?",
|
||||
"最喜欢的运动项目是什么?",
|
||||
"有没有特别支持的球队或运动员?"
|
||||
],
|
||||
[
|
||||
"你会说其他语言吗?",
|
||||
"最想学的语言是什么?",
|
||||
"学习语言有什么技巧吗?"
|
||||
],
|
||||
[
|
||||
"你对科技产品感兴趣吗?",
|
||||
"最近有没有关注什么新科技?",
|
||||
"最喜欢的电子产品是什么?"
|
||||
],
|
||||
[
|
||||
"你喜欢喝什么样的饮料?",
|
||||
"有没有自己调饮料的习惯?",
|
||||
"最喜欢的饮品品牌是什么?"
|
||||
],
|
||||
[
|
||||
"你平时用社交媒体吗?",
|
||||
"常用哪些平台?",
|
||||
"在社交媒体上做什么?"
|
||||
],
|
||||
[
|
||||
"你对艺术感兴趣吗?",
|
||||
"最喜欢的艺术家是谁?",
|
||||
"有去过哪些艺术展览?"
|
||||
],
|
||||
[
|
||||
"你喜欢DIY吗?",
|
||||
"平时做些什么手工?",
|
||||
"有没有完成的作品可以分享?"
|
||||
],
|
||||
[
|
||||
"你喜欢种植植物吗?",
|
||||
"有养什么植物?",
|
||||
"最喜欢的植物是什么?"
|
||||
],
|
||||
[
|
||||
"你喜欢拍照吗?",
|
||||
"喜欢拍什么样的照片?",
|
||||
"有没有用什么特别的摄影设备?"
|
||||
],
|
||||
[
|
||||
"你喜欢听播客吗?",
|
||||
"常听哪些主题的播客?",
|
||||
"有没有推荐的播客?"
|
||||
],
|
||||
[
|
||||
"你对历史感兴趣吗?",
|
||||
"最喜欢哪个历史时期?",
|
||||
"有没有特别喜欢的历史人物?"
|
||||
],
|
||||
[
|
||||
"你喜欢画画吗?",
|
||||
"平时画什么类型的画?",
|
||||
"有参加过画展吗?"
|
||||
],
|
||||
[
|
||||
"你喜欢写作吗?",
|
||||
"平时写什么类型的文章?",
|
||||
"有没有发表过作品?"
|
||||
],
|
||||
[
|
||||
"你喜欢钓鱼吗?",
|
||||
"平时去哪里钓鱼?",
|
||||
"有没有钓到过什么大鱼?"
|
||||
],
|
||||
[
|
||||
"你喜欢露营吗?",
|
||||
"平时会去哪里露营?",
|
||||
"有没有什么难忘的露营经历?"
|
||||
],
|
||||
[
|
||||
"你喜欢摄影吗?",
|
||||
"最喜欢拍什么题材?",
|
||||
"有没有特别喜欢的摄影师?"
|
||||
],
|
||||
[
|
||||
"你喜欢喝酒吗?",
|
||||
"喜欢什么类型的酒?",
|
||||
"有没有推荐的酒吧或品牌?"
|
||||
],
|
||||
[
|
||||
"你喜欢滑雪吗?",
|
||||
"平时去哪里滑雪?",
|
||||
"有没有什么滑雪技巧分享?"
|
||||
],
|
||||
[
|
||||
"你喜欢海边还是山里?",
|
||||
"最喜欢去哪个地方度假?",
|
||||
"有没有什么特别推荐的景点?"
|
||||
],
|
||||
[
|
||||
"你喜欢参加音乐节吗?",
|
||||
"参加过哪些音乐节?",
|
||||
"最喜欢的音乐节是哪一个?"
|
||||
],
|
||||
[
|
||||
"你喜欢跑步吗?",
|
||||
"平时跑多长距离?",
|
||||
"有没有参加过马拉松?"
|
||||
],
|
||||
[
|
||||
"你喜欢参加聚会吗?",
|
||||
"平时和朋友聚会做什么?",
|
||||
"有没有什么有趣的聚会游戏?"
|
||||
],
|
||||
[
|
||||
"你喜欢收集东西吗?",
|
||||
"收集什么类型的物品?",
|
||||
"有没有什么特别的收藏?"
|
||||
]
|
||||
]
|
||||
}
|
||||
@@ -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
|
||||
```
|
||||
|
||||
@@ -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] <green>{time:YYYY-MM-DD HH:mm:ss}</green> | <level>{level.name[0]}</level> | <level>{message}</level>", 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)
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user