mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-09 03:20:54 +00:00
feat(bench): add HaluMem dataset statistics analyzer and update LLM concurrency control
- Add new analyze_dataset_stats.py script for comprehensive HaluMem dataset analysis - Include statistics for user sessions, dialogues, content lengths and chunk distributions - Replace rate limiting with concurrency control in BaseLLM using semaphore mechanism - Update configuration to use max_concurrency instead of max_rps and rps_window - Modify dialogue formatting to include only user messages in evaluation - Add percentile calculations and detailed content size distribution metrics - Implement session splitting logic based on character length thresholds - Provide per-user statistics and summary tables for dataset analysis - Refactor BaseLLM to use internal _chat_impl and _stream_chat_impl methods - Remove rate limiting locks and timestamps from LLM initialization - Add command-line interface for dataset statistics analysis tool
This commit is contained in:
parent
9b69d96b76
commit
327dc58f70
4 changed files with 596 additions and 90 deletions
547
bench/halumem/analyze_dataset_stats.py
Normal file
547
bench/halumem/analyze_dataset_stats.py
Normal file
|
|
@ -0,0 +1,547 @@
|
|||
"""
|
||||
HaluMem Dataset Statistics Analyzer
|
||||
|
||||
统计 HaluMem 数据集的各项指标:
|
||||
- 每个用户的 session 数量
|
||||
- 每个 session 的对话数量
|
||||
- 每个 session 的对话总长度
|
||||
|
||||
Usage:
|
||||
python bench/halumem/analyze_dataset_stats.py --data_path /path/to/HaluMem-Medium.jsonl
|
||||
"""
|
||||
|
||||
import json
|
||||
import re
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
from loguru import logger
|
||||
|
||||
|
||||
@dataclass
|
||||
class UserStats:
|
||||
"""单个用户的统计数据"""
|
||||
user_name: str
|
||||
uuid: str
|
||||
num_sessions: int
|
||||
dialogues_per_session: list[int] # 每个 session 的对话数量
|
||||
dialogue_lengths_per_session: list[int] # 每个 session 的对话总长度(字符数)
|
||||
num_chunks_after_split: int # 按 5000 字符分割后的 chunk 数量
|
||||
|
||||
|
||||
@dataclass
|
||||
class DatasetStats:
|
||||
"""整体数据集统计"""
|
||||
total_users: int
|
||||
total_sessions: int
|
||||
total_dialogues: int
|
||||
|
||||
avg_sessions_per_user: float
|
||||
avg_dialogues_per_session: float
|
||||
avg_dialogue_length_per_session: float
|
||||
|
||||
# 详细分布
|
||||
sessions_per_user_list: list[int]
|
||||
dialogues_per_session_list: list[int]
|
||||
dialogue_lengths_per_session_list: list[int]
|
||||
|
||||
# Content 统计
|
||||
total_contents: int # 所有对话回合的 content 总数
|
||||
content_sizes: list[int] # 每个 content 的大小(字符数)
|
||||
min_content_size: int
|
||||
max_content_size: int
|
||||
percentiles: dict[str, float] # 分位点统计(全部)
|
||||
|
||||
# 按 role 分类的 Content 统计
|
||||
total_user_contents: int
|
||||
total_assistant_contents: int
|
||||
user_percentiles: dict[str, float] # user 角色的分位点
|
||||
assistant_percentiles: dict[str, float] # assistant 角色的分位点
|
||||
|
||||
# Session 分割统计
|
||||
total_chunks_after_split: int # 按 5000 字符分割后的总 chunk 数
|
||||
chunks_per_user_list: list[int] # 每个用户分割后的 chunk 数量
|
||||
avg_chunks_per_user: float # 平均每个用户的 chunk 数量
|
||||
|
||||
|
||||
class DatasetAnalyzer:
|
||||
"""数据集分析器"""
|
||||
|
||||
def __init__(self, data_path: str):
|
||||
self.data_path = data_path
|
||||
self.user_stats_list: list[UserStats] = []
|
||||
self.all_content_sizes: list[int] = [] # 收集所有 content 的大小
|
||||
self.user_content_sizes: list[int] = [] # user 角色的 content 大小
|
||||
self.assistant_content_sizes: list[int] = [] # assistant 角色的 content 大小
|
||||
|
||||
@staticmethod
|
||||
def extract_user_name(persona_info: str) -> str:
|
||||
"""从 persona_info 中提取用户名"""
|
||||
match = re.search(r"Name:\s*(.*?);", persona_info)
|
||||
if not match:
|
||||
return "Unknown"
|
||||
return match.group(1).strip()
|
||||
|
||||
@staticmethod
|
||||
def calculate_dialogue_length(dialogue: list[dict]) -> int:
|
||||
"""计算对话的总长度(字符数)"""
|
||||
total_length = 0
|
||||
for turn in dialogue:
|
||||
content = turn.get("content", "")
|
||||
total_length += len(content)
|
||||
return total_length
|
||||
|
||||
@staticmethod
|
||||
def split_session_into_chunks(dialogue: list[dict], max_length: int = 5000) -> int:
|
||||
"""
|
||||
将一个 session 按照 max_length 分割成多个 chunks。
|
||||
规则:
|
||||
1. 每次添加 2 个对话回合(user-assistant 对)
|
||||
2. 如果添加后超过 max_length,就开始新的 chunk
|
||||
3. 但是每个 chunk 至少包含 2 个对话回合
|
||||
|
||||
返回分割后的 chunk 数量
|
||||
"""
|
||||
if not dialogue:
|
||||
return 0
|
||||
|
||||
chunks = []
|
||||
current_chunk = []
|
||||
current_length = 0
|
||||
|
||||
# 每次处理 2 个对话回合
|
||||
i = 0
|
||||
while i < len(dialogue):
|
||||
# 取 2 个对话回合(如果不足 2 个,取剩余的)
|
||||
pair = dialogue[i:i+2]
|
||||
pair_length = sum(len(turn.get("content", "")) for turn in pair)
|
||||
|
||||
# 如果当前 chunk 为空,直接添加(保证至少 2 个)
|
||||
if not current_chunk:
|
||||
current_chunk.extend(pair)
|
||||
current_length += pair_length
|
||||
i += len(pair)
|
||||
else:
|
||||
# 如果添加这一对后会超过限制
|
||||
if current_length + pair_length > max_length:
|
||||
# 保存当前 chunk,开始新的 chunk
|
||||
chunks.append(current_chunk)
|
||||
current_chunk = pair
|
||||
current_length = pair_length
|
||||
i += len(pair)
|
||||
else:
|
||||
# 否则添加到当前 chunk
|
||||
current_chunk.extend(pair)
|
||||
current_length += pair_length
|
||||
i += len(pair)
|
||||
|
||||
# 添加最后一个 chunk
|
||||
if current_chunk:
|
||||
chunks.append(current_chunk)
|
||||
|
||||
return len(chunks)
|
||||
|
||||
def load_and_analyze(self):
|
||||
"""加载并分析数据集"""
|
||||
logger.info(f"Loading data from: {self.data_path}")
|
||||
|
||||
with open(self.data_path, "r", encoding="utf-8") as f:
|
||||
for line_num, line in enumerate(f, 1):
|
||||
if not line.strip():
|
||||
continue
|
||||
|
||||
try:
|
||||
user_data = json.loads(line)
|
||||
self._analyze_user(user_data)
|
||||
except json.JSONDecodeError as e:
|
||||
logger.error(f"Error parsing line {line_num}: {e}")
|
||||
continue
|
||||
|
||||
logger.info(f"Analyzed {len(self.user_stats_list)} users")
|
||||
|
||||
def _analyze_user(self, user_data: dict):
|
||||
"""分析单个用户的数据"""
|
||||
user_name = self.extract_user_name(user_data.get("persona_info", ""))
|
||||
uuid = user_data.get("uuid", "")
|
||||
sessions = user_data.get("sessions", [])
|
||||
|
||||
dialogues_per_session = []
|
||||
dialogue_lengths_per_session = []
|
||||
total_chunks = 0
|
||||
|
||||
for session in sessions:
|
||||
dialogue = session.get("dialogue", [])
|
||||
num_dialogues = len(dialogue)
|
||||
dialogue_length = self.calculate_dialogue_length(dialogue)
|
||||
|
||||
dialogues_per_session.append(num_dialogues)
|
||||
dialogue_lengths_per_session.append(dialogue_length)
|
||||
|
||||
# 计算这个 session 分割后的 chunk 数量
|
||||
num_chunks = self.split_session_into_chunks(dialogue, max_length=5000)
|
||||
total_chunks += num_chunks
|
||||
|
||||
# 收集每个 content 的大小,并按 role 分类
|
||||
for turn in dialogue:
|
||||
content = turn.get("content", "")
|
||||
content_size = len(content)
|
||||
role = turn.get("role", "")
|
||||
|
||||
self.all_content_sizes.append(content_size)
|
||||
|
||||
if role == "user":
|
||||
self.user_content_sizes.append(content_size)
|
||||
elif role == "assistant":
|
||||
self.assistant_content_sizes.append(content_size)
|
||||
|
||||
user_stats = UserStats(
|
||||
user_name=user_name,
|
||||
uuid=uuid,
|
||||
num_sessions=len(sessions),
|
||||
dialogues_per_session=dialogues_per_session,
|
||||
dialogue_lengths_per_session=dialogue_lengths_per_session,
|
||||
num_chunks_after_split=total_chunks
|
||||
)
|
||||
|
||||
self.user_stats_list.append(user_stats)
|
||||
|
||||
def compute_dataset_stats(self) -> DatasetStats:
|
||||
"""计算整体数据集统计"""
|
||||
total_users = len(self.user_stats_list)
|
||||
|
||||
sessions_per_user_list = [u.num_sessions for u in self.user_stats_list]
|
||||
total_sessions = sum(sessions_per_user_list)
|
||||
|
||||
dialogues_per_session_list = []
|
||||
dialogue_lengths_per_session_list = []
|
||||
|
||||
for user in self.user_stats_list:
|
||||
dialogues_per_session_list.extend(user.dialogues_per_session)
|
||||
dialogue_lengths_per_session_list.extend(user.dialogue_lengths_per_session)
|
||||
|
||||
total_dialogues = sum(dialogues_per_session_list)
|
||||
|
||||
# 计算平均值
|
||||
avg_sessions_per_user = total_sessions / total_users if total_users > 0 else 0
|
||||
avg_dialogues_per_session = (
|
||||
total_dialogues / total_sessions if total_sessions > 0 else 0
|
||||
)
|
||||
avg_dialogue_length_per_session = (
|
||||
sum(dialogue_lengths_per_session_list) / len(dialogue_lengths_per_session_list)
|
||||
if dialogue_lengths_per_session_list else 0
|
||||
)
|
||||
|
||||
# Content 统计
|
||||
total_contents = len(self.all_content_sizes)
|
||||
min_content_size = min(self.all_content_sizes) if self.all_content_sizes else 0
|
||||
max_content_size = max(self.all_content_sizes) if self.all_content_sizes else 0
|
||||
|
||||
# 计算分位点 (10%, 15%, 20%, ..., 95%)
|
||||
percentile_points = list(range(10, 100, 5)) # 10, 15, 20, ..., 95
|
||||
|
||||
# 全部 content 的分位点
|
||||
percentiles = {}
|
||||
if self.all_content_sizes:
|
||||
content_array = np.array(self.all_content_sizes)
|
||||
for p in percentile_points:
|
||||
percentiles[f"p{p}"] = float(np.percentile(content_array, p))
|
||||
|
||||
# user 角色的分位点
|
||||
user_percentiles = {}
|
||||
if self.user_content_sizes:
|
||||
user_array = np.array(self.user_content_sizes)
|
||||
for p in percentile_points:
|
||||
user_percentiles[f"p{p}"] = float(np.percentile(user_array, p))
|
||||
|
||||
# assistant 角色的分位点
|
||||
assistant_percentiles = {}
|
||||
if self.assistant_content_sizes:
|
||||
assistant_array = np.array(self.assistant_content_sizes)
|
||||
for p in percentile_points:
|
||||
assistant_percentiles[f"p{p}"] = float(np.percentile(assistant_array, p))
|
||||
|
||||
# Session 分割统计
|
||||
chunks_per_user_list = [u.num_chunks_after_split for u in self.user_stats_list]
|
||||
total_chunks_after_split = sum(chunks_per_user_list)
|
||||
avg_chunks_per_user = (
|
||||
total_chunks_after_split / total_users if total_users > 0 else 0
|
||||
)
|
||||
|
||||
return DatasetStats(
|
||||
total_users=total_users,
|
||||
total_sessions=total_sessions,
|
||||
total_dialogues=total_dialogues,
|
||||
avg_sessions_per_user=avg_sessions_per_user,
|
||||
avg_dialogues_per_session=avg_dialogues_per_session,
|
||||
avg_dialogue_length_per_session=avg_dialogue_length_per_session,
|
||||
sessions_per_user_list=sessions_per_user_list,
|
||||
dialogues_per_session_list=dialogues_per_session_list,
|
||||
dialogue_lengths_per_session_list=dialogue_lengths_per_session_list,
|
||||
total_contents=total_contents,
|
||||
content_sizes=self.all_content_sizes,
|
||||
min_content_size=min_content_size,
|
||||
max_content_size=max_content_size,
|
||||
percentiles=percentiles,
|
||||
total_user_contents=len(self.user_content_sizes),
|
||||
total_assistant_contents=len(self.assistant_content_sizes),
|
||||
user_percentiles=user_percentiles,
|
||||
assistant_percentiles=assistant_percentiles,
|
||||
total_chunks_after_split=total_chunks_after_split,
|
||||
chunks_per_user_list=chunks_per_user_list,
|
||||
avg_chunks_per_user=avg_chunks_per_user
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _print_percentiles(percentiles: dict[str, float]):
|
||||
"""打印分位点统计(辅助函数)"""
|
||||
if not percentiles:
|
||||
print(" (无数据)")
|
||||
return
|
||||
|
||||
sorted_percentiles = sorted(percentiles.keys(), key=lambda x: int(x[1:]))
|
||||
|
||||
# 每行显示 5 个分位点,让输出更紧凑
|
||||
for i in range(0, len(sorted_percentiles), 5):
|
||||
line_items = []
|
||||
for percentile_key in sorted_percentiles[i:i+5]:
|
||||
percentile_value = percentiles[percentile_key]
|
||||
p_num = percentile_key[1:] # 去掉 'p' 前缀
|
||||
line_items.append(f"{p_num}%: {percentile_value:.0f}")
|
||||
print(f" {' | '.join(line_items)}")
|
||||
|
||||
def print_summary(self, stats: DatasetStats):
|
||||
"""打印统计摘要"""
|
||||
print("\n" + "=" * 80)
|
||||
print("HALUMEM DATASET STATISTICS")
|
||||
print("=" * 80 + "\n")
|
||||
|
||||
print("📊 总体统计:")
|
||||
print(f" 总用户数: {stats.total_users}")
|
||||
print(f" 总 Session 数: {stats.total_sessions}")
|
||||
print(f" 总对话数: {stats.total_dialogues}")
|
||||
|
||||
print(f"\n📈 平均值:")
|
||||
print(f" 每个用户的平均 Session 数: {stats.avg_sessions_per_user:.2f}")
|
||||
print(f" 每个 Session 的平均对话数: {stats.avg_dialogues_per_session:.2f}")
|
||||
print(f" 每个 Session 的平均对话长度(字符): {stats.avg_dialogue_length_per_session:.2f}")
|
||||
|
||||
print(f"\n📊 分布统计:")
|
||||
if stats.sessions_per_user_list:
|
||||
print(f" 每用户 Session 数 - 最小: {min(stats.sessions_per_user_list)}, "
|
||||
f"最大: {max(stats.sessions_per_user_list)}")
|
||||
|
||||
if stats.dialogues_per_session_list:
|
||||
print(f" 每 Session 对话数 - 最小: {min(stats.dialogues_per_session_list)}, "
|
||||
f"最大: {max(stats.dialogues_per_session_list)}")
|
||||
|
||||
if stats.dialogue_lengths_per_session_list:
|
||||
print(f" 每 Session 对话长度 - 最小: {min(stats.dialogue_lengths_per_session_list)}, "
|
||||
f"最大: {max(stats.dialogue_lengths_per_session_list)}")
|
||||
|
||||
print(f"\n💬 Content 详细统计:")
|
||||
print(f" 总 Content 数量: {stats.total_contents}")
|
||||
print(f" User 消息数: {stats.total_user_contents}")
|
||||
print(f" Assistant 消息数: {stats.total_assistant_contents}")
|
||||
print(f" Content 大小(字符数):")
|
||||
print(f" 最小值: {stats.min_content_size}")
|
||||
print(f" 最大值: {stats.max_content_size}")
|
||||
|
||||
if stats.content_sizes:
|
||||
avg_content_size = sum(stats.content_sizes) / len(stats.content_sizes)
|
||||
print(f" 平均值: {avg_content_size:.2f}")
|
||||
|
||||
print(f"\n📈 Content 大小分位点 (全部):")
|
||||
self._print_percentiles(stats.percentiles)
|
||||
|
||||
print(f"\n📈 Content 大小分位点 (User 角色):")
|
||||
self._print_percentiles(stats.user_percentiles)
|
||||
|
||||
print(f"\n📈 Content 大小分位点 (Assistant 角色):")
|
||||
self._print_percentiles(stats.assistant_percentiles)
|
||||
|
||||
print(f"\n✂️ Session 分割统计 (按 5000 字符分割):")
|
||||
print(f" 原始 Session 总数: {stats.total_sessions}")
|
||||
print(f" 分割后 Chunk 总数: {stats.total_chunks_after_split}")
|
||||
print(f" 每个用户平均 Chunk 数: {stats.avg_chunks_per_user:.2f}")
|
||||
print(f" Chunk/Session 比例: {stats.total_chunks_after_split / stats.total_sessions:.2f}x")
|
||||
|
||||
print("\n" + "=" * 80)
|
||||
|
||||
def print_per_user_stats(self):
|
||||
"""打印每个用户的详细统计"""
|
||||
print("\n" + "=" * 80)
|
||||
print("PER-USER STATISTICS")
|
||||
print("=" * 80 + "\n")
|
||||
|
||||
for idx, user_stats in enumerate(self.user_stats_list, 1):
|
||||
avg_dialogues = (
|
||||
sum(user_stats.dialogues_per_session) / len(user_stats.dialogues_per_session)
|
||||
if user_stats.dialogues_per_session else 0
|
||||
)
|
||||
avg_length = (
|
||||
sum(user_stats.dialogue_lengths_per_session) / len(user_stats.dialogue_lengths_per_session)
|
||||
if user_stats.dialogue_lengths_per_session else 0
|
||||
)
|
||||
|
||||
print(f"[{idx}] {user_stats.user_name} (UUID: {user_stats.uuid[:8]}...)")
|
||||
print(f" Session 数: {user_stats.num_sessions}")
|
||||
print(f" 分割后 Chunk 数: {user_stats.num_chunks_after_split}")
|
||||
print(f" 平均每 Session 对话数: {avg_dialogues:.2f}")
|
||||
print(f" 平均每 Session 对话长度: {avg_length:.2f} 字符")
|
||||
print()
|
||||
|
||||
def print_user_split_summary(self):
|
||||
"""打印每个用户的分割统计摘要(表格形式)"""
|
||||
print("\n" + "=" * 80)
|
||||
print("PER-USER SESSION SPLIT SUMMARY (按 5000 字符分割)")
|
||||
print("=" * 80 + "\n")
|
||||
|
||||
# 表头
|
||||
print(f"{'序号':<6} {'用户名':<25} {'原始Sessions':<15} {'分割后Chunks':<15} {'比例':<10}")
|
||||
print("-" * 80)
|
||||
|
||||
# 每个用户的数据
|
||||
for idx, user_stats in enumerate(self.user_stats_list, 1):
|
||||
ratio = (
|
||||
user_stats.num_chunks_after_split / user_stats.num_sessions
|
||||
if user_stats.num_sessions > 0 else 0
|
||||
)
|
||||
print(f"{idx:<6} {user_stats.user_name[:24]:<25} {user_stats.num_sessions:<15} "
|
||||
f"{user_stats.num_chunks_after_split:<15} {ratio:.2f}x")
|
||||
|
||||
print("-" * 80)
|
||||
|
||||
# 总计
|
||||
total_sessions = sum(u.num_sessions for u in self.user_stats_list)
|
||||
total_chunks = sum(u.num_chunks_after_split for u in self.user_stats_list)
|
||||
overall_ratio = total_chunks / total_sessions if total_sessions > 0 else 0
|
||||
|
||||
print(f"{'总计':<6} {'':<25} {total_sessions:<15} {total_chunks:<15} {overall_ratio:.2f}x")
|
||||
print("=" * 80)
|
||||
|
||||
def save_results(self, output_path: str, stats: DatasetStats):
|
||||
"""保存统计结果到 JSON 文件"""
|
||||
results = {
|
||||
"summary": {
|
||||
"total_users": stats.total_users,
|
||||
"total_sessions": stats.total_sessions,
|
||||
"total_dialogues": stats.total_dialogues,
|
||||
"avg_sessions_per_user": stats.avg_sessions_per_user,
|
||||
"avg_dialogues_per_session": stats.avg_dialogues_per_session,
|
||||
"avg_dialogue_length_per_session": stats.avg_dialogue_length_per_session,
|
||||
"session_split_stats": {
|
||||
"total_chunks_after_split": stats.total_chunks_after_split,
|
||||
"avg_chunks_per_user": stats.avg_chunks_per_user,
|
||||
"chunk_to_session_ratio": (
|
||||
stats.total_chunks_after_split / stats.total_sessions
|
||||
if stats.total_sessions > 0 else 0
|
||||
)
|
||||
},
|
||||
"content_stats": {
|
||||
"total_contents": stats.total_contents,
|
||||
"total_user_contents": stats.total_user_contents,
|
||||
"total_assistant_contents": stats.total_assistant_contents,
|
||||
"min_content_size": stats.min_content_size,
|
||||
"max_content_size": stats.max_content_size,
|
||||
"avg_content_size": (
|
||||
sum(stats.content_sizes) / len(stats.content_sizes)
|
||||
if stats.content_sizes else 0
|
||||
),
|
||||
"percentiles_all": stats.percentiles,
|
||||
"percentiles_user": stats.user_percentiles,
|
||||
"percentiles_assistant": stats.assistant_percentiles
|
||||
}
|
||||
},
|
||||
"per_user_stats": [
|
||||
{
|
||||
"user_name": u.user_name,
|
||||
"uuid": u.uuid,
|
||||
"num_sessions": u.num_sessions,
|
||||
"num_chunks_after_split": u.num_chunks_after_split,
|
||||
"avg_dialogues_per_session": (
|
||||
sum(u.dialogues_per_session) / len(u.dialogues_per_session)
|
||||
if u.dialogues_per_session else 0
|
||||
),
|
||||
"avg_dialogue_length_per_session": (
|
||||
sum(u.dialogue_lengths_per_session) / len(u.dialogue_lengths_per_session)
|
||||
if u.dialogue_lengths_per_session else 0
|
||||
),
|
||||
"dialogues_per_session": u.dialogues_per_session,
|
||||
"dialogue_lengths_per_session": u.dialogue_lengths_per_session
|
||||
}
|
||||
for u in self.user_stats_list
|
||||
]
|
||||
}
|
||||
|
||||
with open(output_path, "w", encoding="utf-8") as f:
|
||||
json.dump(results, f, ensure_ascii=False, indent=2)
|
||||
|
||||
logger.info(f"Results saved to: {output_path}")
|
||||
|
||||
|
||||
def main(data_path: str, output_path: str = None, show_per_user: bool = False):
|
||||
"""主函数"""
|
||||
# 检查文件是否存在
|
||||
if not Path(data_path).exists():
|
||||
logger.error(f"File not found: {data_path}")
|
||||
return
|
||||
|
||||
# 创建分析器并执行分析
|
||||
analyzer = DatasetAnalyzer(data_path)
|
||||
analyzer.load_and_analyze()
|
||||
|
||||
# 计算统计数据
|
||||
stats = analyzer.compute_dataset_stats()
|
||||
|
||||
# 打印摘要
|
||||
analyzer.print_summary(stats)
|
||||
|
||||
# 打印每个用户的分割统计摘要(始终显示)
|
||||
analyzer.print_user_split_summary()
|
||||
|
||||
# 打印每个用户的详细统计(可选)
|
||||
if show_per_user:
|
||||
analyzer.print_per_user_stats()
|
||||
|
||||
# 保存结果到文件
|
||||
if output_path:
|
||||
analyzer.save_results(output_path, stats)
|
||||
else:
|
||||
# 默认保存到与数据文件相同目录
|
||||
default_output = str(Path(data_path).parent / "dataset_statistics.json")
|
||||
analyzer.save_results(default_output, stats)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import argparse
|
||||
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Analyze HaluMem dataset statistics"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--data_path",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Path to HaluMem JSONL file"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output_path",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Path to save statistics JSON (default: dataset_statistics.json in same dir)"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--show_per_user",
|
||||
action="store_true",
|
||||
help="Show detailed statistics for each user"
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
main(
|
||||
data_path=args.data_path,
|
||||
output_path=args.output_path,
|
||||
show_per_user=args.show_per_user
|
||||
)
|
||||
|
|
@ -77,15 +77,19 @@ class DataLoader:
|
|||
|
||||
@staticmethod
|
||||
def format_dialogue_for_eval(dialogue: list[dict], user_name: str = None) -> str:
|
||||
"""Format dialogue into string for evaluation."""
|
||||
"""Format dialogue into string for evaluation (only user messages)."""
|
||||
formatted_turns = []
|
||||
for turn in dialogue:
|
||||
# Skip assistant messages - only include user messages
|
||||
if turn['role'] != 'user':
|
||||
continue
|
||||
|
||||
timestamp = datetime.strptime(
|
||||
turn["timestamp"], "%b %d, %Y, %H:%M:%S"
|
||||
).replace(tzinfo=timezone.utc).strftime("%Y-%m-%d %H:%M:%S")
|
||||
|
||||
# Use user_name if role is 'user' and user_name is provided
|
||||
role = user_name if turn['role'] == 'user' and user_name else turn['role']
|
||||
# Use user_name if provided
|
||||
role = user_name if user_name else 'user'
|
||||
|
||||
formatted_turns.append(
|
||||
f"Role: {role}\n"
|
||||
|
|
|
|||
|
|
@ -16,16 +16,13 @@ llm:
|
|||
default:
|
||||
backend: openai
|
||||
model_name: qwen3-30b-a3b-instruct-2507
|
||||
max_rps: 6
|
||||
rps_window: 10
|
||||
max_concurrency: 20
|
||||
|
||||
qwen3_max_instruct:
|
||||
backend: openai
|
||||
model_name: qwen3-max
|
||||
# temperature: 0.6
|
||||
max_rps: 2
|
||||
# max_rps: 9
|
||||
rps_window: 1
|
||||
max_concurrency: 20
|
||||
|
||||
embedding_model:
|
||||
default:
|
||||
|
|
|
|||
|
|
@ -4,7 +4,6 @@ import asyncio
|
|||
import json
|
||||
import time
|
||||
from abc import ABC, abstractmethod
|
||||
from collections import deque
|
||||
from typing import Callable, Generator, AsyncGenerator, Any
|
||||
|
||||
from loguru import logger
|
||||
|
|
@ -18,88 +17,24 @@ from ..schema import ToolCall
|
|||
class BaseLLM(ABC):
|
||||
"""Abstract base class defining the standard interface for LLM interactions."""
|
||||
|
||||
def __init__(self, model_name: str, max_retries: int = 10, raise_exception: bool = False, max_rps: int | None = None, rps_window: float = 1.0, **kwargs):
|
||||
def __init__(self, model_name: str, max_retries: int = 10, raise_exception: bool = False, max_concurrency: int | None = None, **kwargs):
|
||||
"""Initialize the LLM client with model configurations and retry policies.
|
||||
|
||||
Args:
|
||||
model_name: The name of the model to use
|
||||
max_retries: Maximum number of retry attempts on failure
|
||||
raise_exception: Whether to raise exceptions or return default values
|
||||
max_rps: Maximum requests allowed within the time window. If None, no rate limiting is applied.
|
||||
rps_window: Time window in seconds for rate limiting (default: 1.0).
|
||||
For example: max_rps=10, rps_window=5.0 means max 10 requests in 5 seconds.
|
||||
max_concurrency: Maximum concurrent requests for async operations. If None, no concurrency limit is applied.
|
||||
**kwargs: Additional model-specific parameters
|
||||
"""
|
||||
self.model_name: str = model_name
|
||||
self.max_retries: int = max_retries
|
||||
self.raise_exception: bool = raise_exception
|
||||
self.max_rps: int | None = max_rps
|
||||
self.rps_window: float = rps_window
|
||||
self.max_concurrency: int | None = max_concurrency
|
||||
self.kwargs: dict = kwargs
|
||||
|
||||
# Rate limiting state - using deque for efficient O(1) operations
|
||||
self._request_timestamps: deque = deque()
|
||||
self._rate_limit_lock = asyncio.Lock() # For async rate limiting
|
||||
import threading
|
||||
self._rate_limit_lock_sync = threading.Lock() # For sync rate limiting
|
||||
|
||||
async def _wait_for_rate_limit(self):
|
||||
"""Async rate limiting: wait if necessary to respect max_rps constraint within the time window."""
|
||||
if self.max_rps is None:
|
||||
return
|
||||
|
||||
while True:
|
||||
current_time = time.time()
|
||||
|
||||
# Clean up old timestamps and calculate wait time in one critical section
|
||||
async with self._rate_limit_lock:
|
||||
# Remove timestamps older than the time window
|
||||
while self._request_timestamps and current_time - self._request_timestamps[0] >= self.rps_window:
|
||||
self._request_timestamps.popleft()
|
||||
|
||||
# If we have space in the rate limit window, record and proceed immediately
|
||||
if len(self._request_timestamps) < self.max_rps:
|
||||
self._request_timestamps.append(current_time)
|
||||
return
|
||||
|
||||
# Calculate how long to wait until the oldest request expires
|
||||
oldest_timestamp = self._request_timestamps[0]
|
||||
wait_time = oldest_timestamp + self.rps_window - current_time
|
||||
|
||||
# Wait OUTSIDE the lock so other requests can proceed
|
||||
if wait_time > 0:
|
||||
logger.debug(f"Rate limit reached ({self.max_rps} requests in {self.rps_window}s). Waiting {wait_time:.3f}s")
|
||||
await asyncio.sleep(wait_time + 0.001) # Add small buffer to ensure timestamp expires
|
||||
# Loop back to check again after waiting
|
||||
|
||||
def _wait_for_rate_limit_sync(self):
|
||||
"""Synchronous rate limiting: wait if necessary to respect max_rps constraint within the time window."""
|
||||
if self.max_rps is None:
|
||||
return
|
||||
|
||||
while True:
|
||||
current_time = time.time()
|
||||
|
||||
# Clean up old timestamps and calculate wait time in one critical section
|
||||
with self._rate_limit_lock_sync:
|
||||
# Remove timestamps older than the time window
|
||||
while self._request_timestamps and current_time - self._request_timestamps[0] >= self.rps_window:
|
||||
self._request_timestamps.popleft()
|
||||
|
||||
# If we have space in the rate limit window, record and proceed immediately
|
||||
if len(self._request_timestamps) < self.max_rps:
|
||||
self._request_timestamps.append(current_time)
|
||||
return
|
||||
|
||||
# Calculate how long to wait until the oldest request expires
|
||||
oldest_timestamp = self._request_timestamps[0]
|
||||
wait_time = oldest_timestamp + self.rps_window - current_time
|
||||
|
||||
# Wait OUTSIDE the lock so other requests can proceed
|
||||
if wait_time > 0:
|
||||
logger.debug(f"Rate limit reached ({self.max_rps} requests in {self.rps_window}s). Waiting {wait_time:.3f}s")
|
||||
time.sleep(wait_time + 0.001) # Add small buffer to ensure timestamp expires
|
||||
# Loop back to check again after waiting
|
||||
# Concurrency control for async operations
|
||||
self._semaphore: asyncio.Semaphore | None = asyncio.Semaphore(max_concurrency) if max_concurrency else None
|
||||
|
||||
@staticmethod
|
||||
def _accumulate_tool_call_chunk(tool_call, ret_tools: list[ToolCall]):
|
||||
|
|
@ -191,9 +126,23 @@ class BaseLLM(ABC):
|
|||
model_name: Optional model name to override self.model_name
|
||||
**kwargs: Additional parameters
|
||||
"""
|
||||
# Apply rate limiting before making the request
|
||||
await self._wait_for_rate_limit()
|
||||
|
||||
# Apply concurrency control if configured
|
||||
if self._semaphore:
|
||||
async with self._semaphore:
|
||||
async for chunk in self._stream_chat_impl(messages, tools, model_name, **kwargs):
|
||||
yield chunk
|
||||
else:
|
||||
async for chunk in self._stream_chat_impl(messages, tools, model_name, **kwargs):
|
||||
yield chunk
|
||||
|
||||
async def _stream_chat_impl(
|
||||
self,
|
||||
messages: list[Message],
|
||||
tools: list[ToolCall] | None = None,
|
||||
model_name: str | None = None,
|
||||
**kwargs,
|
||||
) -> AsyncGenerator[StreamChunk, None]:
|
||||
"""Internal implementation of stream_chat with retry logic."""
|
||||
stream_kwargs = self._build_stream_kwargs(messages, tools, model_name=model_name, **kwargs)
|
||||
|
||||
for i in range(self.max_retries):
|
||||
|
|
@ -229,9 +178,6 @@ class BaseLLM(ABC):
|
|||
model_name: Optional model name to override self.model_name
|
||||
**kwargs: Additional parameters
|
||||
"""
|
||||
# Apply rate limiting before making the request
|
||||
self._wait_for_rate_limit_sync()
|
||||
|
||||
stream_kwargs = self._build_stream_kwargs(messages, tools, model_name=model_name, **kwargs)
|
||||
|
||||
for i in range(self.max_retries):
|
||||
|
|
@ -408,14 +354,29 @@ class BaseLLM(ABC):
|
|||
model_name: Optional model name to override self.model_name
|
||||
**kwargs: Additional parameters
|
||||
"""
|
||||
# Apply concurrency control if configured
|
||||
if self._semaphore:
|
||||
async with self._semaphore:
|
||||
return await self._chat_impl(messages, tools, enable_stream_print, callback_fn, default_value, model_name, **kwargs)
|
||||
else:
|
||||
return await self._chat_impl(messages, tools, enable_stream_print, callback_fn, default_value, model_name, **kwargs)
|
||||
|
||||
async def _chat_impl(
|
||||
self,
|
||||
messages: list[Message],
|
||||
tools: list[ToolCall] | None = None,
|
||||
enable_stream_print: bool = False,
|
||||
callback_fn: Callable[[Message], Any] | None = None,
|
||||
default_value: Any = None,
|
||||
model_name: str | None = None,
|
||||
**kwargs,
|
||||
) -> Message | Any:
|
||||
"""Internal implementation of chat with retry and error handling logic."""
|
||||
# Use the provided model_name or fall back to self.model_name
|
||||
effective_model = model_name if model_name is not None else self.model_name
|
||||
|
||||
for i in range(self.max_retries):
|
||||
try:
|
||||
# Apply rate limiting before making the request
|
||||
await self._wait_for_rate_limit()
|
||||
|
||||
result = await self._chat(
|
||||
messages=messages,
|
||||
tools=tools,
|
||||
|
|
@ -489,9 +450,6 @@ class BaseLLM(ABC):
|
|||
|
||||
for i in range(self.max_retries):
|
||||
try:
|
||||
# Apply rate limiting before making the request
|
||||
self._wait_for_rate_limit_sync()
|
||||
|
||||
result = self._chat_sync(
|
||||
messages=messages,
|
||||
tools=tools,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue