优化数据处理逻辑,重构消息合并功能。

This commit is contained in:
xming521
2025-04-09 22:27:06 +08:00
parent 360ff7332d
commit 019925f601
3 changed files with 401 additions and 87 deletions
+86 -86
View File
@@ -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,
},
}
-1
View File
@@ -49,7 +49,6 @@
"图片"
],
"history_length": 10,
"prefer_comma": false, //聊天时是不是喜欢使用逗号,
"conversation_strategy": "time_window", // 基于时间窗口的判断策略
"time_window": 10, // 时间窗口(分钟),
"prompt_with_history": false // 是否在prompt中包含历史对话
+315
View File
@@ -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