zsy814's picture
Initial EchoLoc code release
d8bfe4a verified
Raw
History Blame Contribute Delete
31.6 kB
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
import argparse
import base64
import json
import os
import re
import time
from collections import defaultdict
from concurrent.futures import ThreadPoolExecutor, as_completed
from pathlib import Path
from typing import Any, Dict, List, Tuple
from openai import OpenAI
from tqdm import tqdm
try:
import matplotlib.pyplot as plt
except ImportError:
plt = None
try:
import pandas as pd
except ImportError:
pd = None
if plt is not None:
plt.rcParams.update({"font.size": 10})
SUB_CATS = {
"acoustic_attributes": [
"acoustic_attributes/age",
"acoustic_attributes/speed",
"acoustic_attributes/gender",
"acoustic_attributes/emotion",
"acoustic_attributes/pitch",
"acoustic_attributes/volume",
"acoustic_attributes/composite_properties",
],
"instruction": [
"instruction/emotion",
"instruction/variation",
"instruction/style",
],
"role_play": [
"role_play/character",
"role_play/scenario",
],
"empathy": [
"empathy/anger",
"empathy/sadness_disappointment",
"empathy/anxiety_fear",
"empathy/joy_excitement",
],
}
# --------------------------- Utility Functions --------------------------------
def load_prompts(prompts_dir: Path) -> Dict[str, str]:
mapping = {}
for file in prompts_dir.glob("*.txt"):
mapping[file.stem] = file.read_text(encoding="utf-8")
if not mapping:
raise RuntimeError(f"No *.txt prompt files found in {prompts_dir}")
return mapping
def construct_prompt(template: str, instruction: str, ability: str) -> str:
return template.format(instruction_type=ability, input_instruction=instruction)
def debug_print_input(sample, root_dir, prompts):
"""跑之前调用这个,肉眼确认输入"""
big_cat = sample["ability"].split("/")[0]
prompt_text = construct_prompt(
prompts[big_cat],
instruction=sample.get("instruct_text", ""),
ability=sample["ability"],
)
audio_path = Path(sample.get("final_audio_path") or root_dir / sample.get("response_audio_path", ""))
print("=" * 60)
print(f"[id] {sample['id']}")
print(f"[ability] {sample['ability']}")
print(f"[instruct_text]{sample['instruct_text']}")
print(f"[audio_path] {audio_path} exists={audio_path.exists()}")
print(f"[prompt末尾50字]{prompt_text[-200:]}")
print("=" * 60)
def ensure_dir(p: Path):
p.mkdir(parents=True, exist_ok=True)
def safe_float(x: Any):
try:
return float(x)
except (TypeError, ValueError):
return None
def normalize_score_str(x: Any) -> str:
v = safe_float(x)
if v is None:
return ""
if 1 <= v <= 5:
if float(v).is_integer():
return str(int(v))
return str(v).rstrip("0").rstrip(".")
return ""
def parse_score(model_reply: str) -> str:
if not model_reply:
return ""
text = str(model_reply).strip()
# 1) 优先匹配 [[4]] / [[4.0]]
matches = re.findall(r"\[\[\s*([0-9]+(?:\.[0-9]+)?)\s*]]", text, flags=re.S)
for s in reversed(matches):
score = normalize_score_str(s)
if score:
return score
# 2) 整体就是数字
score = normalize_score_str(text)
if score:
return score
# 3) 只匹配明确的最终分数字段,避免把 0-100 的风格强度误读成 1-5 分
matches = re.findall(
r"(?:final\s+score|score|rating|gemini_score|分数|评分|最终分数)"
r"\s*[::=]?\s*\[?\[?\s*([1-5](?:\.[0-9]+)?)(?!\d)\s*\]?\]?",
text,
flags=re.I,
)
for s in reversed(matches):
score = normalize_score_str(s)
if score:
return score
return ""
def encode_audio_base64(audio_path: Path) -> str:
with audio_path.open("rb") as f:
return base64.b64encode(f.read()).decode("utf-8")
def build_retry_prompt(original_prompt: str, attempt: int) -> str:
return f"""
{original_prompt}
IMPORTANT FORMAT REQUIREMENT:
Your previous response did not contain a valid score.
You MUST now output exactly ONE score between 1 and 5.
Do not explain.
Do not output any extra words.
Preferred format: [[score]]
Valid examples:
[[1]]
[[2]]
[[3]]
[[4]]
[[5]]
This is retry attempt #{attempt}.
""".strip()
def normalize_id(value: Any):
if value is None:
return None
try:
return int(value)
except (TypeError, ValueError):
return str(value)
def load_dataset_metadata(dataset_metadata_path: Path) -> Dict[Any, Dict[str, Any]]:
dataset_map: Dict[Any, Dict[str, Any]] = {}
if not dataset_metadata_path or not Path(dataset_metadata_path).exists():
return dataset_map
with Path(dataset_metadata_path).open(encoding="utf-8") as f:
for line in f:
try:
d = json.loads(line)
except json.JSONDecodeError:
continue
instruct_id = d.get("instruct_id", d.get("id"))
item = {
"id": d.get("id", instruct_id),
"instruct_id": instruct_id,
"ability": d.get("ability", ""),
"instruct_text": d.get("instruct_text") or d.get("audio_content", ""),
"response_audio_path": d.get("response_audio_path") or d.get("file_name", ""),
}
for key in (d.get("id"), instruct_id):
norm = normalize_id(key)
if norm is not None:
dataset_map[norm] = item
return dataset_map
def candidate_path(root_dir: Path, path_str: Any) -> Path:
if not path_str:
return None
path = Path(str(path_str).strip())
if path.is_absolute():
return path
return root_dir / path
def resolve_audio_path(root_dir: Path, sample: Dict[str, Any], dataset_entry: Dict[str, Any] = None) -> str:
iid = sample.get("instruct_id", sample.get("id"))
candidates = []
for key in ("final_audio_path", "response_audio_path", "file_name", "response_wav"):
path = candidate_path(root_dir, sample.get(key))
if path:
candidates.append(path)
if dataset_entry:
path = candidate_path(root_dir, dataset_entry.get("response_audio_path"))
if path:
candidates.append(path)
if iid is not None:
candidates.append(root_dir / f"{iid}.wav")
seen = set()
for path in candidates:
key = str(path)
if key in seen:
continue
seen.add(key)
if path.exists():
return str(path)
return ""
# --------------------------- Proxy Gemini (OpenAI-compatible) -----------------
def call_proxy_gemini_api(
prompt_text: str,
response_audio_path: Path,
model_name: str,
max_retry: int,
sleep_between_retry: int,
api_key: str,
base_url: str,
temperature: float = 1.0,
top_p: float = 0.7,
max_tokens: int = 4096,
verbose: bool = False,
) -> Tuple[bool, str, str]:
client = OpenAI(api_key=api_key, base_url=base_url)
last_err_msg = ""
last_reply_text = ""
audio_b64 = encode_audio_base64(response_audio_path)
audio_format = response_audio_path.suffix.replace(".", "").lower() or "wav"
for attempt in range(1, max_retry + 1):
try:
current_prompt = prompt_text if attempt == 1 else build_retry_prompt(prompt_text, attempt)
content = [{"type": "text", "text": current_prompt}]
content.append(
{
"type": "input_audio",
"input_audio": {
"data": audio_b64,
"format": audio_format,
},
}
)
resp = client.chat.completions.create(
model=model_name,
messages=[{"role": "user", "content": content}],
max_tokens=max_tokens,
temperature=temperature,
top_p=top_p,
)
if not hasattr(resp, "choices") or not resp.choices:
last_err_msg = f"Proxy API bad response: {resp}"
print(f"[Attempt {attempt}/{max_retry}] bad response: {last_err_msg}")
else:
reply_text = resp.choices[0].message.content or ""
last_reply_text = reply_text
score_str = parse_score(reply_text)
if verbose:
print(f"[Attempt {attempt}/{max_retry}] raw reply: {reply_text[:200]}")
if score_str:
return True, reply_text, score_str
last_err_msg = f"Unable to parse score from reply: {reply_text[:200]}"
print(f"[Attempt {attempt}/{max_retry}] parse failed")
except Exception as e:
last_err_msg = str(e)
print(f"[Attempt {attempt}/{max_retry}] API call error: {last_err_msg}")
if attempt < max_retry:
time.sleep(sleep_between_retry)
return False, (last_reply_text or last_err_msg), ""
# --------------------------- Evaluation Core --------------------------------
def evaluate_one(
sample: dict,
root_dir: Path,
prompts: Dict[str, str],
out_dir: Path,
model_name: str,
max_retry: int,
sleep_between_retry: int,
api_key: str,
base_url: str,
fallback_score: str,
temperature: float,
top_p: float,
max_tokens: int,
verbose: bool,
) -> dict:
ability = sample["ability"]
big_cat, small_cat = ability.split("/", 1)
sample_model_name = sample.get("model_name", "model_eval")
sample_id = sample.get("id") or sample.get("instruct_id") or "unknown_id"
sample["id"] = sample_id
if big_cat not in prompts:
error_msg = f"Prompt template not found for category {big_cat}"
sample["gemini_score"] = fallback_score
sample["gemini_status"] = "fallback_prompt_missing"
sample["gemini_error"] = error_msg
sample["gemini_raw"] = error_msg
return sample
response_audio_path_str = sample.get("final_audio_path", "")
if not response_audio_path_str:
error_msg = f"Audio file not found for ID {sample_id} in any expected location."
sample["gemini_score"] = fallback_score
sample["gemini_status"] = "fallback_audio_not_found"
sample["gemini_error"] = error_msg
sample["gemini_raw"] = error_msg
return sample
response_audio_path = Path(response_audio_path_str) # 此时已经是探测好的绝对路径了
if not response_audio_path.exists():
error_msg = f"Audio file not found: {response_audio_path}"
sample["gemini_score"] = fallback_score
sample["gemini_status"] = "fallback_audio_not_found"
sample["gemini_error"] = error_msg
sample["gemini_raw"] = error_msg
return sample
prompt_text = construct_prompt(
prompts[big_cat],
instruction=sample.get("audio_content", "") or sample.get("instruct_text", ""),
ability=sample["ability"],
)
success, reply, score = call_proxy_gemini_api(
prompt_text=prompt_text,
response_audio_path=response_audio_path,
model_name=model_name,
max_retry=min(max_retry, 5),
sleep_between_retry=sleep_between_retry,
api_key=api_key,
base_url=base_url,
temperature=temperature,
top_p=top_p,
max_tokens=max_tokens,
verbose=verbose,
)
if success:
sample["gemini_score"] = score
sample["gemini_raw"] = reply
sample["gemini_status"] = "success"
else:
sample["gemini_score"] = fallback_score
sample["gemini_raw"] = reply
sample["gemini_status"] = "fallback_after_5_failures"
sample["gemini_error"] = f"Failed after {min(max_retry, 5)} attempts"
# ← debug 日志必须放在这里,不能放在上面
debug_path = out_dir / "debug_inputs_outputs.txt"
with debug_path.open("a", encoding="utf-8") as dbg:
dbg.write("=" * 80 + "\n")
dbg.write(f"[ID] {sample_id}\n")
dbg.write(f"[ABILITY] {ability}\n")
dbg.write(f"[AUDIO] {response_audio_path}\n")
dbg.write(f"[INSTRUCT] {sample.get('instruct_text', '')}\n")
dbg.write(f"[PROMPT]\n{prompt_text}\n")
dbg.write(f"[RAW OUTPUT]\n{sample.get('gemini_raw', '')}\n")
dbg.write(f"[SCORE] {sample.get('gemini_score', '')}\n")
dbg.write(f"[STATUS] {sample.get('gemini_status', '')}\n")
dbg.write("=" * 80 + "\n\n")
save_path = (
out_dir
/ sample_model_name
/ "gemini_proxy_res"
/ big_cat
/ small_cat
/ f"{sample_id}.txt"
)
ensure_dir(save_path.parent)
save_path.write_text(sample.get("gemini_raw", ""), encoding="utf-8")
return sample
def load_processed_ids_and_clean_failures(scored_jsonl: Path, overwrite: bool) -> Tuple[set, int]:
processed = set()
cleaned_count = 0
if not scored_jsonl.exists() or overwrite:
return processed, cleaned_count
valid_lines = []
total_lines = 0
with scored_jsonl.open("r", encoding="utf-8") as f:
for line in f:
total_lines += 1
line = line.strip()
if not line:
continue
try:
rec = json.loads(line)
gemini_score = rec.get("gemini_score", "")
if safe_float(gemini_score) is not None:
processed.add(normalize_id(rec.get("id")))
valid_lines.append(line)
else:
cleaned_count += 1
print(f"🗑️ Removing failed record ID {rec.get('id')}: {gemini_score}")
except Exception as e:
cleaned_count += 1
print(f"🗑️ Removing invalid JSON line: {str(e)}")
if cleaned_count > 0:
print(f"📝 Cleaning {cleaned_count} failed records from {scored_jsonl}")
with scored_jsonl.open("w", encoding="utf-8") as f:
for line in valid_lines:
f.write(line + "\n")
print(f"✅ Cleaned file saved. Kept {len(valid_lines)} valid records out of {total_lines} total lines.")
return processed, cleaned_count
def remove_duplicates_from_tasks(tasks: List[dict]) -> Tuple[List[dict], Dict[str, int]]:
seen_ids = set()
unique_tasks = []
duplicate_count = 0
duplicate_details = defaultdict(int)
for task in tasks:
task_id = normalize_id(task.get("id"))
ability = task.get("ability", "unknown")
if task_id not in seen_ids:
seen_ids.add(task_id)
unique_tasks.append(task)
else:
duplicate_count += 1
duplicate_details[ability] += 1
print(f"🔄 Removing duplicate sample ID: {task_id} (ability: {ability})")
stats = {
"total_duplicates": duplicate_count,
"by_ability": dict(duplicate_details),
}
return unique_tasks, stats
def eval_all_samples(
root_dir: Path,
metadata_path: Path,
prompts_dir: Path,
out_dir: Path,
max_per_ability: int,
concurrency: int,
overwrite: bool,
model_name: str,
max_retry: int,
sleep_between_retry: int,
api_key: str,
base_url: str,
fallback_score: str,
dataset_metadata_path: Path = None,
temperature: float = 1.0,
top_p: float = 0.7,
max_tokens: int = 4096,
verbose: bool = False,
) -> Path:
prompts = load_prompts(prompts_dir)
dataset_map = load_dataset_metadata(dataset_metadata_path) if dataset_metadata_path else {}
if dataset_map:
print(f"📚 Loaded {len(dataset_map)} entries from dataset: {dataset_metadata_path}")
scored_jsonl = out_dir / "metadata_with_score.jsonl"
processed_ids, _ = load_processed_ids_and_clean_failures(scored_jsonl, overwrite)
tasks: List[dict] = []
per_ability_counter = defaultdict(int)
skipped_missing_dataset = 0
skipped_missing_required = 0
print(f"📖 Reading samples from: {metadata_path}")
with metadata_path.open(encoding="utf-8") as fin:
for line in fin:
try:
sample = json.loads(line)
except json.JSONDecodeError:
continue
iid = sample.get("instruct_id", sample.get("id"))
iid_norm = normalize_id(iid)
dataset_entry = dataset_map.get(iid_norm) if dataset_map else None
# 传入官方 dataset metadata 时,用官方文本/类别覆盖 result.jsonl,
# 这样评分严格对齐 VStyle 的 ground-truth instruction。
if dataset_map:
if not dataset_entry:
skipped_missing_dataset += 1
continue
entry = dataset_entry
sample["id"] = entry["id"]
sample["instruct_id"] = entry["instruct_id"]
sample["ability"] = entry["ability"]
sample["instruct_text"] = entry["instruct_text"]
else:
sample["id"] = sample.get("id", iid)
sample["instruct_id"] = sample.get("instruct_id", iid)
sample["instruct_text"] = sample.get("instruct_text") or sample.get("audio_content", "")
if not sample.get("id") or not sample.get("ability") or not sample.get("instruct_text"):
skipped_missing_required += 1
continue
sample["final_audio_path"] = resolve_audio_path(root_dir, sample, dataset_entry)
if not str(sample.get("model_name", "")).strip():
sample["model_name"] = "unknown_model"
ability = sample["ability"]
if normalize_id(sample["id"]) in processed_ids:
continue
if per_ability_counter[ability] >= max_per_ability:
continue
per_ability_counter[ability] += 1
tasks.append(sample)
print(f"📊 Step 1 - After initial filtering: {len(tasks)} samples")
if skipped_missing_dataset:
print(f"📊 Skipped {skipped_missing_dataset} samples not found in dataset metadata")
if skipped_missing_required:
print(f"📊 Skipped {skipped_missing_required} samples missing id/ability/instruction")
print(f"📊 Step 2 - Before deduplication: {len(tasks)} samples")
tasks, dup_stats = remove_duplicates_from_tasks(tasks)
print(f"📊 Step 3 - After deduplication: {len(tasks)} samples (removed {dup_stats['total_duplicates']} duplicates)")
# for s in tasks[:3]:
# debug_print_input(s, root_dir, prompts)
# input("按回车继续全量评测...")
if not tasks:
print("\n✅ No new samples to process, program ended.")
return scored_jsonl
ensure_dir(scored_jsonl.parent)
write_mode = "a" if scored_jsonl.exists() and not overwrite else "w"
print(f"\n🎯 Starting evaluation of {len(tasks)} samples...")
print("=" * 80)
with scored_jsonl.open(write_mode, encoding="utf-8") as fout, \
ThreadPoolExecutor(max_workers=concurrency) as executor, \
tqdm(total=len(tasks), desc="Evaluating") as pbar:
futures = {
executor.submit(
evaluate_one,
s,
root_dir,
prompts,
out_dir,
model_name,
max_retry,
sleep_between_retry,
api_key,
base_url,
fallback_score,
temperature,
top_p,
max_tokens,
verbose,
): s
for s in tasks
}
for fut in as_completed(futures):
res = fut.result()
fout.write(json.dumps(res, ensure_ascii=False) + "\n")
fout.flush()
pbar.update(1)
print(f"\n✅ All processing complete, wrote {len(tasks)} new results -> {scored_jsonl}")
return scored_jsonl
# --------------------------- Analysis and Visualization ------------------------------
def get_ordered_abilities():
ordered_abilities = []
for big_cat in ["acoustic_attributes", "instruction", "role_play", "empathy"]:
ordered_abilities.extend(SUB_CATS[big_cat])
return ordered_abilities
def analyze_scores(scored_jsonl: Path, out_dir: Path):
if pd is None or plt is None:
missing = []
if pd is None:
missing.append("pandas")
if plt is None:
missing.append("matplotlib")
print(f"⚠️ Skipping score analysis/plots because missing dependency: {', '.join(missing)}")
print(" Install them if you need summary plots: pip install pandas matplotlib")
return
records = [
json.loads(l)
for l in scored_jsonl.read_text(encoding="utf-8").splitlines()
if l.strip()
]
rows, parse_fail_ids = [], []
for rec in records:
score = safe_float(rec.get("gemini_score"))
if score is None:
parse_fail_ids.append(rec.get("id"))
continue
rows.append(
{
"id": rec.get("id"),
"model": rec.get("model_name"),
"ability": rec.get("ability"),
"score": score,
}
)
print("len(rows):", len(rows))
if not rows:
print("❌ No usable score data, script terminated.")
return
df = pd.DataFrame(rows)
BIG_CAT_W = 0.25
WEIGHTS = {}
ap_comp = "acoustic_attributes/composite_properties"
ap_comp_w = BIG_CAT_W * 0.5
ap_rem_each = (BIG_CAT_W - ap_comp_w) / (len(SUB_CATS["acoustic_attributes"]) - 1)
for ab in SUB_CATS["acoustic_attributes"]:
WEIGHTS[ab] = ap_comp_w if ab == ap_comp else ap_rem_each
for big in ("instruction", "role_play", "empathy"):
each = BIG_CAT_W / len(SUB_CATS[big])
for ab in SUB_CATS[big]:
WEIGHTS[ab] = each
ordered_abilities = get_ordered_abilities()
print("\n================= Statistical Results =================")
for model, g in df.groupby("model"):
print(f"\nModel: {model}")
sub_means = g.groupby("ability")["score"].mean().to_dict()
weighted_sum, used_w = 0.0, 0.0
for ab, w in WEIGHTS.items():
if ab in sub_means:
weighted_sum += sub_means[ab] * w
used_w += w
weighted_mean = weighted_sum / used_w if used_w else float("nan")
print(f" Weighted overall average score: {weighted_mean:.2f} (weight coverage={used_w:.1%})")
print(" Major category average scores:")
for big_cat in ["acoustic_attributes", "instruction", "role_play", "empathy"]:
big_cat_weighted_sum = 0.0
big_cat_used_w = 0.0
big_cat_sample_count = 0
for ability in SUB_CATS[big_cat]:
if ability in sub_means:
weight = WEIGHTS[ability]
big_cat_weighted_sum += sub_means[ability] * weight
big_cat_used_w += weight
big_cat_sample_count += g[g["ability"] == ability].shape[0]
if big_cat_used_w > 0:
big_cat_avg = big_cat_weighted_sum / big_cat_used_w
coverage = big_cat_used_w / BIG_CAT_W
print(f" {big_cat:<20s}: {big_cat_avg:.2f} (n={big_cat_sample_count}, coverage={coverage:.1%})")
else:
print(f" {big_cat:<20s}: N/A (no data)")
print(" Ability average scores:")
ability_means = g.groupby("ability")["score"].mean().to_dict()
for ab in ordered_abilities:
if ab in ability_means:
cnt = g[g["ability"] == ab].shape[0]
weight = WEIGHTS.get(ab, 0.0)
print(f" {ab:<40s}: {ability_means[ab]:.2f} (n={cnt}, weight={weight:.3f})")
print("\nStarting visualization...")
for model, g in df.groupby("model"):
model_dir = out_dir / model
ensure_dir(model_dir)
sub_means = g.groupby("ability")["score"].mean().to_dict()
big_cat_scores = []
big_cat_labels = []
for big_cat in ["acoustic_attributes", "instruction", "role_play", "empathy"]:
big_cat_weighted_sum = 0.0
big_cat_used_w = 0.0
for ability in SUB_CATS[big_cat]:
if ability in sub_means:
weight = WEIGHTS[ability]
big_cat_weighted_sum += sub_means[ability] * weight
big_cat_used_w += weight
if big_cat_used_w > 0:
big_cat_avg = big_cat_weighted_sum / big_cat_used_w
big_cat_scores.append(big_cat_avg)
big_cat_labels.append(big_cat.replace("_", " ").title())
if big_cat_scores:
plt.figure(figsize=(10, 6))
bars = plt.bar(range(len(big_cat_scores)), big_cat_scores)
plt.xticks(range(len(big_cat_labels)), big_cat_labels, rotation=45, ha="right")
for bar, score in zip(bars, big_cat_scores):
plt.text(
bar.get_x() + bar.get_width() / 2,
bar.get_height() + 0.01,
f"{score:.2f}",
ha="center",
va="bottom",
)
w_sum = sum(sub_means[ab] * w for ab, w in WEIGHTS.items() if ab in sub_means)
w_used = sum(w for ab, w in WEIGHTS.items() if ab in sub_means)
w_avg = w_sum / w_used if w_used else float("nan")
plt.axhline(w_avg, linestyle="--", linewidth=2, label=f"Overall average: {w_avg:.2f}")
plt.ylabel("Average Score")
plt.title(f"{model} - Major Category Scores")
plt.legend()
plt.grid(True, alpha=0.3)
plt.tight_layout()
plt.savefig(model_dir / "major_category_scores.png", dpi=200, bbox_inches="tight")
plt.close()
ability_mean = g.groupby("ability")["score"].mean()
ordered_data, ordered_labels = [], []
for ab in ordered_abilities:
if ab in ability_mean.index:
ordered_data.append(ability_mean[ab])
ordered_labels.append(ab)
plt.figure(figsize=(8, 4 + 0.25 * len(ordered_data)))
y_pos = range(len(ordered_data))
plt.barh(y_pos, ordered_data)
plt.yticks(y_pos, ordered_labels)
sub_means = g.groupby("ability")["score"].mean().to_dict()
w_sum = sum(sub_means[ab] * w for ab, w in WEIGHTS.items() if ab in sub_means)
w_used = sum(w for ab, w in WEIGHTS.items() if ab in sub_means)
w_avg = w_sum / w_used if w_used else float("nan")
plt.axvline(w_avg, linestyle="--", label="Weighted overall average")
plt.xlabel("Average score")
plt.title(f"{model} - Ability average scores")
plt.legend()
plt.tight_layout()
plt.savefig(model_dir / "ability_scores.png", dpi=200)
plt.close()
print(f"✅ Visualization complete, saved to folder: {out_dir}")
print("\n================= Parse Failure Information =================")
if parse_fail_ids:
print(f"Total of {len(parse_fail_ids)} records failed gemini_score parsing.")
else:
print("All records' gemini_score parsed successfully.")
# ----------------------------- CLI ----------------------------------
def main():
root = "./data/examples"
parser = argparse.ArgumentParser(
description="Concurrent proxy scoring for voice results with forced score retries"
)
parser.add_argument("--root_dir", default=f"{root}/model_res/en/wav", help="wav directory path")
parser.add_argument("--metadata_path", default=f"{root}/model_res/en/metadata.jsonl", help="metadata.jsonl path")
parser.add_argument("--prompts_dir", default="lalm_eval/eval_prompts/en", help="Directory containing prompt templates")
parser.add_argument("--out_dir", default=f"{root}/eval_res/en", help="Output root directory")
parser.add_argument("--temperature", type=float, default=1.0)
parser.add_argument("--top_p", type=float, default=0.7)
parser.add_argument("--max_tokens", type=int, default=4096)
parser.add_argument(
"--api_key",
default=os.environ.get("VSTYLE_API_KEY") or os.environ.get("OPENAI_API_KEY"),
help="Proxy API Key; defaults to VSTYLE_API_KEY or OPENAI_API_KEY",
)
parser.add_argument(
"--base_url",
default=os.environ.get("VSTYLE_BASE_URL"),
help="Proxy base_url, e.g. https://xxx/v1; defaults to VSTYLE_BASE_URL",
)
parser.add_argument(
"--model_name",
default=os.environ.get("VSTYLE_JUDGE_MODEL", "gemini-2.5-pro"),
help="Model name on proxy; defaults to VSTYLE_JUDGE_MODEL or gemini-2.5-pro",
)
parser.add_argument(
"--max_retry_api",
type=int,
default=5,
help="Maximum number of retry attempts, capped at 5 in scoring logic",
)
parser.add_argument(
"--sleep_between_retry",
type=int,
default=5,
help="Sleep time between retries (seconds)",
)
parser.add_argument(
"--max_per_ability",
type=int,
default=100000,
help="Maximum number of evaluations per ability",
)
parser.add_argument("--concurrency", type=int, default=4, help="Number of concurrent threads")
parser.add_argument("--overwrite", action="store_true", help="Overwrite existing results")
parser.add_argument(
"--fallback_score",
type=str,
default="3",
help="Fallback score if 5 attempts all fail; should be between 1 and 5",
)
parser.add_argument(
"--dataset_metadata_path",
type=str,
default=None,
help="原始 VStyle 数据集的 metadata.jsonl 路径,用于补充 instruct_text",
)
parser.add_argument("--verbose", action="store_true", help="Print raw judge replies during scoring")
args = parser.parse_args()
if not args.api_key:
parser.error("--api_key is required, or set VSTYLE_API_KEY / OPENAI_API_KEY")
if not args.base_url:
parser.error("--base_url is required, or set VSTYLE_BASE_URL")
if normalize_score_str(args.fallback_score) == "":
raise ValueError("--fallback_score must be a number between 1 and 5")
args.fallback_score = normalize_score_str(args.fallback_score)
print(f"args: {args}")
out_dir = Path(args.out_dir)
ensure_dir(out_dir)
scored_jsonl = eval_all_samples(
root_dir=Path(args.root_dir),
metadata_path=Path(args.metadata_path),
prompts_dir=Path(args.prompts_dir),
out_dir=out_dir,
max_per_ability=args.max_per_ability,
concurrency=args.concurrency,
overwrite=args.overwrite,
model_name=args.model_name,
max_retry=args.max_retry_api,
sleep_between_retry=args.sleep_between_retry,
api_key=args.api_key,
base_url=args.base_url,
fallback_score=args.fallback_score,
dataset_metadata_path=args.dataset_metadata_path,
temperature=args.temperature,
top_p=args.top_p,
max_tokens=args.max_tokens,
verbose=args.verbose,
)
analyze_scores(scored_jsonl, out_dir)
if __name__ == "__main__":
main()