mirror of
https://github.com/xming521/WeClone.git
synced 2026-08-28 18:07:28 +08:00
优化数据处理逻辑,重构消息合并功能。
This commit is contained in:
@@ -21,6 +21,18 @@ class DataProcessor:
|
||||
self.data = None
|
||||
self.processed_data = 1
|
||||
self.csv_folder = "./data/csv"
|
||||
self.type_list = [
|
||||
"文本",
|
||||
"图片",
|
||||
"卡片式链接",
|
||||
"合并转发的聊天记录",
|
||||
"视频",
|
||||
"语言",
|
||||
"未知",
|
||||
"分享的小程序",
|
||||
]
|
||||
self.skip_type_list = self.type_list.copy()
|
||||
self.skip_type_list.remove("文本")
|
||||
# 根据self.config.make_dataset_args.conversation_strategy 判断初始化哪一个策略类
|
||||
if self.config["conversation_strategy"] == "time_window":
|
||||
self.conversation_strategy = TimeWindowStrategy(
|
||||
@@ -47,7 +59,7 @@ class DataProcessor:
|
||||
for csv_file in csv_files:
|
||||
chat_messages = self.load_csv(csv_file)
|
||||
# 第一次预处理后 将chat_message 加入rcsv_df_list
|
||||
message_list.append(self.group_consecutive_messages(chat_messages))
|
||||
message_list.append(self.group_consecutive_messages(messages=chat_messages))
|
||||
# self.process_by_msgtype(chat_message)
|
||||
|
||||
def group_consecutive_messages(
|
||||
@@ -65,68 +77,24 @@ class DataProcessor:
|
||||
if not messages:
|
||||
return []
|
||||
|
||||
grouped_messages = []
|
||||
current_group = [messages[0]]
|
||||
def _combine_messages(messages: List[ChatMessage]) -> ChatMessage:
|
||||
"""
|
||||
合并多条消息为一条
|
||||
|
||||
for i in range(1, len(messages)):
|
||||
current_msg = messages[i]
|
||||
last_msg = current_group[-1]
|
||||
Args:
|
||||
messages: 要合并的消息列表
|
||||
|
||||
# 判断是否是同一个人的连续消息
|
||||
if (
|
||||
current_msg.is_sender == last_msg.is_sender
|
||||
and current_msg.talker == last_msg.talker
|
||||
and (current_msg.CreateTime - last_msg.CreateTime).total_seconds()
|
||||
< self.config.get("message_time_window", 3600)
|
||||
):
|
||||
Returns:
|
||||
ChatMessage: 合并后的消息
|
||||
"""
|
||||
base_msg = messages[0]
|
||||
combined_content = messages[0].msg.strip()
|
||||
|
||||
# 同一个人的连续消息,添加到当前组
|
||||
current_group.append(current_msg)
|
||||
else:
|
||||
# 不是同一个人的消息,处理当前组并开始新组
|
||||
if len(current_group) > 1:
|
||||
# 合并消息内容
|
||||
combined_msg = self._combine_messages(current_group)
|
||||
grouped_messages.append(combined_msg)
|
||||
else:
|
||||
# 只有一条消息,直接添加
|
||||
grouped_messages.append(current_group[0])
|
||||
for i in messages[1:]:
|
||||
content = i.msg.strip()
|
||||
if not content:
|
||||
continue
|
||||
|
||||
# 开始新组
|
||||
current_group = [current_msg]
|
||||
|
||||
# 处理最后一组消息
|
||||
if current_group:
|
||||
if len(current_group) > 1:
|
||||
combined_msg = self._combine_messages(current_group)
|
||||
grouped_messages.append(combined_msg)
|
||||
else:
|
||||
grouped_messages.append(current_group[0])
|
||||
|
||||
return grouped_messages
|
||||
|
||||
def _combine_messages(self, messages: List[ChatMessage]) -> ChatMessage:
|
||||
"""
|
||||
合并多条消息为一条
|
||||
|
||||
Args:
|
||||
messages: 要合并的消息列表
|
||||
|
||||
Returns:
|
||||
ChatMessage: 合并后的消息
|
||||
"""
|
||||
# 以第一条消息为基础
|
||||
base_msg = messages[0]
|
||||
|
||||
# 合并消息内容
|
||||
combined_content = ""
|
||||
for i, msg in enumerate(messages):
|
||||
content = msg.msg.strip()
|
||||
if not content:
|
||||
continue
|
||||
|
||||
if i > 0:
|
||||
# 如果前一条消息没有以标点符号结尾,添加一个句号
|
||||
if combined_content and combined_content[-1] not in [
|
||||
"。",
|
||||
"!",
|
||||
@@ -137,32 +105,64 @@ class DataProcessor:
|
||||
]:
|
||||
combined_content += ","
|
||||
|
||||
combined_content += content
|
||||
combined_content += content
|
||||
|
||||
# 确保最后一条消息以句号结尾
|
||||
if combined_content and combined_content[-1] not in [
|
||||
"。",
|
||||
"!",
|
||||
"?",
|
||||
"…",
|
||||
".",
|
||||
]:
|
||||
combined_content += "。"
|
||||
combined_message = ChatMessage(
|
||||
id=base_msg.id,
|
||||
MsgSvrID=base_msg.MsgSvrID,
|
||||
type_name=base_msg.type_name,
|
||||
is_sender=base_msg.is_sender,
|
||||
talker=base_msg.talker,
|
||||
room_name=base_msg.room_name,
|
||||
msg=combined_content,
|
||||
src=base_msg.src,
|
||||
CreateTime=messages[-1].CreateTime, # 使用最后一条消息的时间
|
||||
)
|
||||
|
||||
# 创建新的合并消息
|
||||
combined_message = ChatMessage(
|
||||
id=base_msg.id,
|
||||
MsgSvrID=base_msg.MsgSvrID,
|
||||
type_name=base_msg.type_name,
|
||||
is_sender=base_msg.is_sender,
|
||||
talker=base_msg.talker,
|
||||
room_name=base_msg.room_name,
|
||||
msg=combined_content,
|
||||
src=base_msg.src,
|
||||
CreateTime=messages[-1].CreateTime, # 使用最后一条消息的时间
|
||||
)
|
||||
return combined_message
|
||||
|
||||
return combined_message
|
||||
grouped_messages = []
|
||||
current_group = []
|
||||
|
||||
for _, current_msg in enumerate(messages):
|
||||
|
||||
if current_msg.type_name in self.skip_type_list:
|
||||
continue
|
||||
|
||||
if not current_group:
|
||||
current_group = [current_msg]
|
||||
continue
|
||||
|
||||
last_msg = current_group[-1]
|
||||
|
||||
# 判断是否是同一个人的连续消息
|
||||
if (
|
||||
current_msg.is_sender == last_msg.is_sender
|
||||
and current_msg.talker == last_msg.talker
|
||||
and (current_msg.CreateTime - last_msg.CreateTime).total_seconds()
|
||||
< 3600
|
||||
):
|
||||
current_group.append(current_msg)
|
||||
else:
|
||||
# 不是同一个人的消息,处理当前组并开始新组
|
||||
if len(current_group) > 1:
|
||||
combined_msg = _combine_messages(current_group)
|
||||
grouped_messages.append(combined_msg)
|
||||
else:
|
||||
grouped_messages.append(current_group[0])
|
||||
|
||||
# 开始新组
|
||||
current_group = [current_msg]
|
||||
|
||||
# 处理最后一组消息
|
||||
if current_group:
|
||||
if len(current_group) > 1:
|
||||
combined_msg = _combine_messages(current_group)
|
||||
grouped_messages.append(combined_msg)
|
||||
else:
|
||||
grouped_messages.append(current_group[0])
|
||||
|
||||
return grouped_messages
|
||||
|
||||
def create_conversation_data(self, messages: List[ChatMessage]) -> dict:
|
||||
"""
|
||||
@@ -180,8 +180,8 @@ class DataProcessor:
|
||||
conversation.append(
|
||||
{
|
||||
"role": "user" if msg.is_sender == 0 else "assistant",
|
||||
"content": msg.content,
|
||||
"timestamp": msg.timestamp,
|
||||
"content": msg.msg,
|
||||
"timestamp": msg.CreateTime,
|
||||
}
|
||||
)
|
||||
|
||||
@@ -195,8 +195,8 @@ class DataProcessor:
|
||||
return {
|
||||
"conversation": conversation,
|
||||
"metadata": {
|
||||
"conversation_id": messages[0].conversation_id if messages else None,
|
||||
"timestamp": messages[-1].timestamp if messages else None,
|
||||
"conversation_id": str(messages[0].id) if messages else None,
|
||||
"timestamp": messages[-1].CreateTime if messages else None,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@@ -49,7 +49,6 @@
|
||||
"图片"
|
||||
],
|
||||
"history_length": 10,
|
||||
"prefer_comma": false, //聊天时是不是喜欢使用逗号,
|
||||
"conversation_strategy": "time_window", // 基于时间窗口的判断策略
|
||||
"time_window": 10, // 时间窗口(分钟),
|
||||
"prompt_with_history": false // 是否在prompt中包含历史对话
|
||||
|
||||
@@ -0,0 +1,315 @@
|
||||
import sys
|
||||
import os
|
||||
import pytest
|
||||
from datetime import datetime, timedelta
|
||||
|
||||
# 添加项目根目录到sys.path
|
||||
current_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
root_dir = os.path.dirname(current_dir)
|
||||
sys.path.append(root_dir)
|
||||
|
||||
from make_dataset.models import ChatMessage
|
||||
from make_dataset.qa_generator import DataProcessor
|
||||
|
||||
# 将当前工作目录更改为项目根目录
|
||||
os.chdir(root_dir)
|
||||
|
||||
# # 测试数据处理器类的初始化和配置加载
|
||||
# def test_data_processor_init():
|
||||
# """测试DataProcessor初始化"""
|
||||
# processor = DataProcessor()
|
||||
# assert processor.csv_folder == "./data/csv"
|
||||
# assert "文本" not in processor.skip_type_list
|
||||
# assert len(processor.type_list) == 8
|
||||
|
||||
|
||||
class MockDataProcessor(DataProcessor):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def processor():
|
||||
"""创建一个测试用的处理器实例"""
|
||||
return MockDataProcessor()
|
||||
|
||||
|
||||
def test_empty_messages(processor):
|
||||
"""测试空消息列表的情况"""
|
||||
messages = []
|
||||
result = processor.group_consecutive_messages(messages)
|
||||
assert result == []
|
||||
|
||||
|
||||
def test_single_message(processor):
|
||||
"""测试单条消息的情况"""
|
||||
now = datetime.now()
|
||||
message = ChatMessage(
|
||||
id=1,
|
||||
MsgSvrID=1001,
|
||||
type_name="文本",
|
||||
is_sender=0,
|
||||
talker="user1",
|
||||
room_name="testroom",
|
||||
msg="你好",
|
||||
src="",
|
||||
CreateTime=now,
|
||||
)
|
||||
|
||||
result = processor.group_consecutive_messages([message])
|
||||
assert len(result) == 1
|
||||
assert result[0].msg == "你好"
|
||||
|
||||
|
||||
def test_consecutive_messages_same_sender(processor):
|
||||
"""测试同一发送者的连续消息"""
|
||||
now = datetime.now()
|
||||
messages = [
|
||||
ChatMessage(
|
||||
id=1,
|
||||
MsgSvrID=1001,
|
||||
type_name="文本",
|
||||
is_sender=0,
|
||||
talker="user1",
|
||||
room_name="testroom",
|
||||
msg="你好",
|
||||
src="",
|
||||
CreateTime=now,
|
||||
),
|
||||
ChatMessage(
|
||||
id=2,
|
||||
MsgSvrID=1002,
|
||||
type_name="文本",
|
||||
is_sender=0,
|
||||
talker="user1",
|
||||
room_name="testroom",
|
||||
msg="最近怎么样",
|
||||
src="",
|
||||
CreateTime=now + timedelta(minutes=5),
|
||||
),
|
||||
ChatMessage(
|
||||
id=3,
|
||||
MsgSvrID=1003,
|
||||
type_name="文本",
|
||||
is_sender=0,
|
||||
talker="user1",
|
||||
room_name="testroom",
|
||||
msg="我想问个问题",
|
||||
src="",
|
||||
CreateTime=now + timedelta(minutes=10),
|
||||
),
|
||||
]
|
||||
|
||||
result = processor.group_consecutive_messages(messages)
|
||||
assert len(result) == 1
|
||||
assert result[0].msg == "你好,最近怎么样,我想问个问题"
|
||||
|
||||
|
||||
def test_messages_different_senders(processor):
|
||||
"""测试不同发送者的消息"""
|
||||
now = datetime.now()
|
||||
messages = [
|
||||
ChatMessage(
|
||||
id=1,
|
||||
MsgSvrID=1001,
|
||||
type_name="文本",
|
||||
is_sender=0,
|
||||
talker="user1",
|
||||
room_name="testroom",
|
||||
msg="你好",
|
||||
src="",
|
||||
CreateTime=now,
|
||||
),
|
||||
ChatMessage(
|
||||
id=2,
|
||||
MsgSvrID=1002,
|
||||
type_name="文本",
|
||||
is_sender=1,
|
||||
talker="user2",
|
||||
room_name="testroom",
|
||||
msg="你好,有什么可以帮你的",
|
||||
src="",
|
||||
CreateTime=now + timedelta(minutes=5),
|
||||
),
|
||||
ChatMessage(
|
||||
id=3,
|
||||
MsgSvrID=1003,
|
||||
type_name="文本",
|
||||
is_sender=0,
|
||||
talker="user1",
|
||||
room_name="testroom",
|
||||
msg="我想问个问题",
|
||||
src="",
|
||||
CreateTime=now + timedelta(minutes=10),
|
||||
),
|
||||
]
|
||||
|
||||
result = processor.group_consecutive_messages(messages)
|
||||
assert len(result) == 3
|
||||
assert result[0].msg == "你好"
|
||||
assert result[1].msg == "你好,有什么可以帮你的"
|
||||
assert result[2].msg == "我想问个问题"
|
||||
|
||||
|
||||
def test_skip_non_text_messages(processor):
|
||||
"""测试跳过非文本消息"""
|
||||
now = datetime.now()
|
||||
messages = [
|
||||
ChatMessage(
|
||||
id=1,
|
||||
MsgSvrID=1001,
|
||||
type_name="文本",
|
||||
is_sender=0,
|
||||
talker="user1",
|
||||
room_name="testroom",
|
||||
msg="你好",
|
||||
src="",
|
||||
CreateTime=now,
|
||||
),
|
||||
ChatMessage(
|
||||
id=2,
|
||||
MsgSvrID=1002,
|
||||
type_name="图片",
|
||||
is_sender=0,
|
||||
talker="user1",
|
||||
room_name="testroom",
|
||||
msg="",
|
||||
src="image.jpg",
|
||||
CreateTime=now + timedelta(minutes=1),
|
||||
),
|
||||
ChatMessage(
|
||||
id=3,
|
||||
MsgSvrID=1003,
|
||||
type_name="文本",
|
||||
is_sender=0,
|
||||
talker="user1",
|
||||
room_name="testroom",
|
||||
msg="看到图片了吗",
|
||||
src="",
|
||||
CreateTime=now + timedelta(minutes=2),
|
||||
),
|
||||
]
|
||||
|
||||
result = processor.group_consecutive_messages(messages)
|
||||
assert len(result) == 1
|
||||
assert result[0].msg == "你好,看到图片了吗"
|
||||
|
||||
|
||||
def test_time_window_limit(processor):
|
||||
"""测试时间窗口限制(超过1小时的消息不会合并)"""
|
||||
now = datetime.now()
|
||||
messages = [
|
||||
ChatMessage(
|
||||
id=1,
|
||||
MsgSvrID=1001,
|
||||
type_name="文本",
|
||||
is_sender=0,
|
||||
talker="user1",
|
||||
room_name="testroom",
|
||||
msg="你好",
|
||||
src="",
|
||||
CreateTime=now,
|
||||
),
|
||||
ChatMessage(
|
||||
id=2,
|
||||
MsgSvrID=1002,
|
||||
type_name="文本",
|
||||
is_sender=0,
|
||||
talker="user1",
|
||||
room_name="testroom",
|
||||
msg="晚上好",
|
||||
src="",
|
||||
CreateTime=now + timedelta(hours=2), # 超过1小时
|
||||
),
|
||||
]
|
||||
|
||||
result = processor.group_consecutive_messages(messages)
|
||||
assert len(result) == 2
|
||||
assert result[0].msg == "你好"
|
||||
assert result[1].msg == "晚上好"
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import pandas as pd
|
||||
import os
|
||||
from datetime import datetime
|
||||
|
||||
# 创建一个测试函数,将group_consecutive_messages的结果转为DataFrame并保存为CSV
|
||||
def test_save_grouped_messages_to_csv(processor):
|
||||
"""测试将group_consecutive_messages的结果转换为DataFrame并保存为CSV"""
|
||||
now = datetime.now()
|
||||
messages = [
|
||||
ChatMessage(
|
||||
id=1,
|
||||
MsgSvrID=1001,
|
||||
type_name="文本",
|
||||
is_sender=0,
|
||||
talker="user1",
|
||||
room_name="testroom",
|
||||
msg="你好",
|
||||
src="",
|
||||
CreateTime=now,
|
||||
),
|
||||
ChatMessage(
|
||||
id=2,
|
||||
MsgSvrID=1002,
|
||||
type_name="文本",
|
||||
is_sender=0,
|
||||
talker="user1",
|
||||
room_name="testroom",
|
||||
msg="这是测试消息",
|
||||
src="",
|
||||
CreateTime=now + timedelta(minutes=10),
|
||||
),
|
||||
ChatMessage(
|
||||
id=3,
|
||||
MsgSvrID=1003,
|
||||
type_name="文本",
|
||||
is_sender=1,
|
||||
talker="user2",
|
||||
room_name="testroom",
|
||||
msg="收到了",
|
||||
src="",
|
||||
CreateTime=now + timedelta(minutes=20),
|
||||
),
|
||||
]
|
||||
|
||||
# 获取分组后的消息
|
||||
grouped_messages = processor.group_consecutive_messages(messages)
|
||||
|
||||
# 将ChatMessage对象转换为字典列表
|
||||
messages_dict = []
|
||||
for msg in grouped_messages:
|
||||
messages_dict.append({
|
||||
"id": msg.id,
|
||||
"MsgSvrID": msg.MsgSvrID,
|
||||
"type_name": msg.type_name,
|
||||
"is_sender": msg.is_sender,
|
||||
"talker": msg.talker,
|
||||
"room_name": msg.room_name,
|
||||
"msg": msg.msg,
|
||||
"src": msg.src,
|
||||
"CreateTime": msg.CreateTime,
|
||||
})
|
||||
|
||||
# 创建DataFrame
|
||||
df = pd.DataFrame(messages_dict)
|
||||
|
||||
# 确保输出目录存在
|
||||
output_dir = "./test_output"
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
|
||||
# 保存为CSV文件
|
||||
output_file = os.path.join(output_dir, f"grouped_messages_{now.strftime('%Y%m%d_%H%M%S')}.csv")
|
||||
df.to_csv(output_file, index=False, encoding="utf-8")
|
||||
|
||||
# 验证文件已创建
|
||||
assert os.path.exists(output_file)
|
||||
|
||||
# 读取CSV文件并验证内容
|
||||
loaded_df = pd.read_csv(output_file)
|
||||
assert len(loaded_df) == len(grouped_messages)
|
||||
assert loaded_df.iloc[0]["msg"] == "你好,这是测试消息" # 验证第一条消息已合并
|
||||
|
||||
print(f"已成功保存分组消息到: {output_file}")
|
||||
return output_file
|
||||
Reference in New Issue
Block a user