Commit 63589d3c 63589d3c84acde4f51e11b0857525c12c41cbca2 by 沈秋雨

更新支持pg下载多版本歌曲

1 parent 5ae39abf
......@@ -13,10 +13,11 @@
python scripts/acrcloud/generate_acrcloud_testset.py \
--audio-dir /Volumes/移动硬盘/composition_test \
--negative-audio-dir /Volumes/移动硬盘/composition_drop \
--out-dir acrcloud_testset_cloud \
--num-songs 20 \
--num-negative-songs 80 \
--seed 123
--out-dir /Volumes/移动硬盘/acrcloud_testset_cloud \
--num-songs 600 \
--num-negative-songs 399 \
--seed 123 \
--negative-variants
输出:
reference.csv — 参照曲(原始文件),需提前入库
......@@ -79,6 +80,8 @@ POSITIVE_VARIANTS: list[tuple[str, str | None]] = [
("codec_320k", "acodec=libmp3lame,b:a=320k"),
# 小范围音效:轻微混响
("reverb_small", "aecho=0.8:0.88:60:0.4"),
# 10 秒短片段:应保留足够指纹特征,应被识别为去重
("short_clip", None),
]
# --------------------------------------------------------------------------
......@@ -91,8 +94,6 @@ NEGATIVE_VARIANTS: list[tuple[str, str | None]] = [
# 升调 / 降调(改变音高,破坏指纹频域特征)
("pitch_up2", "aresample=22050,asetrate=22050*1.1225,aresample=22050"),
("pitch_down2", "aresample=22050,asetrate=22050*0.8909,aresample=22050"),
# 极端片段(过短的片段不足以提取有效指纹)
("short_clip", None),
]
# 片段拼接负样本:每个样本从几首不同歌曲各取一段拼接
......@@ -295,38 +296,84 @@ def main() -> None:
ref_rows = []
query_rows = []
# 断点续跑:加载已有 CSV,跳过已生成的条目
ref_path = out_dir / "reference.csv"
query_path = out_dir / "queries.csv"
done_ref_ids: set[str] = set()
done_query_paths: set[str] = set()
if ref_path.exists():
try:
with ref_path.open(newline="", encoding="utf-8") as f:
for row in csv.DictReader(f):
done_ref_ids.add(row["song_id"])
logger.info("续跑:已加载 reference.csv,跳过 %d 首参照歌", len(done_ref_ids))
except (UnicodeDecodeError, Exception) as e:
logger.warning("reference.csv 读取失败(%s),将重新生成", e)
if query_path.exists():
try:
with query_path.open(newline="", encoding="utf-8") as f:
for row in csv.DictReader(f):
done_query_paths.add(row["audio_path"])
logger.info("续跑:已加载 queries.csv,跳过 %d 条查询", len(done_query_paths))
except (UnicodeDecodeError, Exception) as e:
logger.warning("queries.csv 读取失败(%s),将重新生成", e)
# ---- 应去重(expected=duplicate):参照歌 + 轻微变换 ----
for wav in _tqdm(selected, desc="生成正样本变体", total=len(selected)):
song_id = _song_id(wav)
ref_rows.append({
"song_id": song_id,
"audio_path": str(wav.resolve()),
"variant": "original",
})
if song_id not in done_ref_ids:
ref_rows.append({
"song_id": song_id,
"audio_path": str(wav.resolve()),
"variant": "original",
})
# 参照歌原始(自身查询,验证入库和识别链路)
query_rows.append({
"song_id": song_id,
"audio_path": str(wav.resolve()),
"variant": "acr_original",
"sample_class": "positive",
"expected_song_id": song_id,
"expected": "duplicate",
})
audio_path_str = str(wav.resolve())
if audio_path_str not in done_query_paths:
query_rows.append({
"song_id": song_id,
"audio_path": audio_path_str,
"variant": "acr_original",
"sample_class": "positive",
"expected_song_id": song_id,
"expected": "duplicate",
})
for variant_name, af in POSITIVE_VARIANTS:
dst = variants_dir / f"{song_id}_{variant_name}.wav"
if variant_name == "slight_trim":
ok = _ffmpeg_trim(wav, dst, start_ratio=0.05, duration_ratio=0.90)
else:
ok = _ffmpeg_variant(wav, dst, af)
if not ok:
logger.warning("正样本变换失败,跳过: %s %s", wav.name, variant_name)
dst_str = str(dst.resolve())
if dst_str in done_query_paths:
continue
if not dst.exists():
if variant_name == "slight_trim":
ok = _ffmpeg_trim(wav, dst, start_ratio=0.05, duration_ratio=0.90)
elif variant_name == "short_clip":
duration = _probe_duration(wav)
if duration is not None:
usable_end = max(0.0, duration - 10.0)
start_min = min(duration * 0.30, usable_end)
start_max = min(duration * 0.50, usable_end)
ss = random.uniform(start_min, start_max) if start_max > start_min else start_min
ok = _run_ffmpeg([
"ffmpeg", "-y", "-i", str(wav),
"-ss", f"{ss:.3f}", "-t", "10.0",
"-ar", "22050", "-ac", "1",
str(dst),
])
else:
ok = False
else:
ok = _ffmpeg_variant(wav, dst, af)
if not ok:
logger.warning("正样本变换失败,跳过: %s %s", wav.name, variant_name)
continue
query_rows.append({
"song_id": song_id,
"audio_path": str(dst.resolve()),
"audio_path": dst_str,
"variant": variant_name,
"sample_class": "positive",
"expected_song_id": song_id,
......@@ -336,9 +383,12 @@ def main() -> None:
# ---- 不应去重(expected=not_duplicate):composition_drop 原始文件 ----
for wav in _tqdm(negative_selected, desc="生成负样本原始", total=len(negative_selected)):
song_id = _song_id(wav)
audio_path_str = str(wav.resolve())
if audio_path_str in done_query_paths:
continue
query_rows.append({
"song_id": song_id,
"audio_path": str(wav.resolve()),
"audio_path": audio_path_str,
"variant": "negative_original",
"sample_class": "negative",
"expected_song_id": "",
......@@ -353,30 +403,17 @@ def main() -> None:
for variant_name, af in NEGATIVE_VARIANTS:
dst = variants_dir / f"{song_id}_{variant_name}.wav"
if variant_name == "short_clip":
# 截取 5 秒极短片段,不足以提取有效指纹
duration = _probe_duration(wav)
if duration is not None:
usable_end = max(0.0, duration - 5.0)
start_min = min(duration * 0.30, usable_end)
start_max = min(duration * 0.50, usable_end)
ss = random.uniform(start_min, start_max) if start_max > start_min else start_min
ok = _run_ffmpeg([
"ffmpeg", "-y", "-i", str(wav),
"-ss", f"{ss:.3f}", "-t", "5.0",
"-ar", "22050", "-ac", "1",
str(dst),
])
else:
ok = False
else:
ok = _ffmpeg_variant(wav, dst, af)
if not ok:
logger.warning("破坏性变换失败,跳过: %s %s", wav.name, variant_name)
dst_str = str(dst.resolve())
if dst_str in done_query_paths:
continue
if not dst.exists():
ok = _ffmpeg_variant(wav, dst, af)
if not ok:
logger.warning("破坏性变换失败,跳过: %s %s", wav.name, variant_name)
continue
query_rows.append({
"song_id": song_id,
"audio_path": str(dst.resolve()),
"audio_path": dst_str,
"variant": variant_name,
"sample_class": "negative",
"expected_song_id": song_id,
......@@ -389,48 +426,58 @@ def main() -> None:
srcs = random.sample(selected, SPLICE_SONGS_PER_SAMPLE)
splice_id = f"splice_{i:04d}"
dst = variants_dir / f"{splice_id}_negative_splice.wav"
ok = _ffmpeg_splice(srcs, dst)
if not ok:
logger.warning("片段拼接生成失败,跳过: splice %d", i)
dst_str = str(dst.resolve())
if dst_str in done_query_paths:
continue
if not dst.exists():
ok = _ffmpeg_splice(srcs, dst)
if not ok:
logger.warning("片段拼接生成失败,跳过: splice %d", i)
continue
query_rows.append({
"song_id": splice_id,
"audio_path": str(dst.resolve()),
"audio_path": dst_str,
"variant": "negative_splice",
"sample_class": "negative",
"expected_song_id": "",
"expected": "not_duplicate",
})
# ---- 输出 CSV ----
ref_path = out_dir / "reference.csv"
query_path = out_dir / "queries.csv"
# ---- 输出 CSV(追加模式,只写本次新增行)----
fieldnames = ["song_id", "audio_path", "variant", "sample_class", "expected_song_id", "expected"]
with ref_path.open("w", newline="", encoding="utf-8") as f:
ref_is_new = not ref_path.exists() or done_ref_ids == set()
with ref_path.open("a" if ref_path.exists() else "w", newline="", encoding="utf-8") as f:
writer = csv.DictWriter(f, fieldnames=["song_id", "audio_path", "variant"])
writer.writeheader()
if ref_is_new:
writer.writeheader()
writer.writerows(ref_rows)
with query_path.open("w", newline="", encoding="utf-8") as f:
query_is_new = not query_path.exists() or done_query_paths == set()
with query_path.open("a" if query_path.exists() else "w", newline="", encoding="utf-8") as f:
writer = csv.DictWriter(f, fieldnames=fieldnames)
writer.writeheader()
if query_is_new:
writer.writeheader()
writer.writerows(query_rows)
total_ref = len(done_ref_ids) + len(ref_rows)
total_query = len(done_query_paths) + len(query_rows)
pos = sum(1 for r in query_rows if r["expected"] == "duplicate")
neg = sum(1 for r in query_rows if r["expected"] == "not_duplicate")
logger.info("参照集: %s (%d 条)", ref_path, len(ref_rows))
logger.info("查询集: %s (%d 条,正样本 %d,负样本 %d)", query_path, len(query_rows), pos, neg)
logger.info("参照集: %s(累计 %d 条,本次新增 %d 条)", ref_path, total_ref, len(ref_rows))
logger.info("查询集: %s(累计 %d 条,本次新增 %d 条,正样本 +%d,负样本 +%d)",
query_path, total_query, len(query_rows), pos, neg)
# 按 sample_class + variant 统计
# 按 sample_class + variant 统计(全量)
from collections import Counter
by_class = Counter(r["sample_class"] for r in query_rows)
with query_path.open(newline="", encoding="utf-8") as f:
all_query_rows = list(csv.DictReader(f))
by_class = Counter(r["sample_class"] for r in all_query_rows)
for cls, cnt in sorted(by_class.items()):
logger.info(" %-20s %d 条", cls, cnt)
by_variant = Counter(r["variant"] for r in query_rows)
by_variant = Counter(r["variant"] for r in all_query_rows)
for variant, cnt in sorted(by_variant.items()):
logger.info(" %-25s %d 条", variant, cnt)
if __name__ == "__main__":
main()
......
......@@ -12,7 +12,8 @@
用法:
conda activate hikoon-data-spider
python scripts/aliyun_dna/evaluate_aliyun_dna.py \
--queries acrcloud_testset_cloud/queries.csv \
--queries /Volumes/移动硬盘/acrcloud_testset_cloud/queries.csv \
--concurrency 8 \
--out results/aliyun_dna_eval.csv
# 只评测特定 variant
......@@ -35,8 +36,10 @@ import logging
import os
import sys
import tempfile
import threading
import time
import urllib.request
from concurrent.futures import ThreadPoolExecutor, as_completed
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parent.parent.parent))
......@@ -229,6 +232,10 @@ def main() -> None:
help=f"降采样目标采样率(默认 {DEFAULT_SAMPLE_RATE}Hz)")
parser.add_argument("--no-resample", action="store_true",
help="不降采样,使用原始文件")
parser.add_argument("--save-interval", type=int, default=20,
help="每处理 N 条就追加保存一次结果(默认 10,0 表示只在结束时保存)")
parser.add_argument("--concurrency", type=int, default=1,
help="并发查询数(默认 1,建议不超过 8 以免触发限流)")
args = parser.parse_args()
# 验证配置
......@@ -256,16 +263,46 @@ def main() -> None:
logger.info("评测样本过滤: 原始 %d 条,保留 %d 条", original_count, len(rows))
oss_bucket = _get_oss_client()
ice_client = _get_ice_client()
out_path = Path(args.out)
out_path.parent.mkdir(parents=True, exist_ok=True)
tmp_files = [] # 跟踪临时降采样文件
fieldnames = ["query_song_id", "audio_song_id", "audio_path", "variant", "sample_class",
"expected_song_id", "expected", "top1_song_id", "top1_similarity",
"top1_hit", "topk_hit", "expected_rank", "expected_similarity",
"expected_duplicate", "predicted_duplicate", "correct",
"upload_ms", "poll_ms", "total_ms", "error"]
result_rows = []
for i, row in enumerate(rows, 1):
# 断点续跑:加载已处理的结果
done_paths: set[str] = set()
result_rows: list[dict] = []
if out_path.exists():
with out_path.open(newline="", encoding="utf-8") as f:
for r in csv.DictReader(f):
result_rows.append(r)
done_paths.add(r["audio_path"])
logger.info("断点续跑:已加载 %d 条历史结果,跳过已处理条目", len(result_rows))
rows = [r for r in rows if r["audio_path"] not in done_paths]
logger.info("本次待处理: %d 条", len(rows))
if not rows:
logger.info("所有样本均已处理完毕")
else:
oss_bucket = _get_oss_client()
ice_client = _get_ice_client()
def _flush_results(rows_to_write: list[dict]) -> None:
with out_path.open("w", newline="", encoding="utf-8") as f:
writer = csv.DictWriter(f, fieldnames=fieldnames)
writer.writeheader()
writer.writerows(rows_to_write)
total_pending = len(rows)
lock = threading.Lock()
completed_count = 0
def _process_one(row: dict, idx: int) -> dict:
"""处理单条查询,tmp 文件用完即删,返回结果 dict(含错误时也返回,不抛出)。"""
audio_path = row["audio_path"]
query_song_id = row.get("song_id") or _song_id_from_audio_path(audio_path)
audio_song_id = _song_id_from_audio_path(audio_path)
......@@ -273,6 +310,7 @@ def main() -> None:
expected_dup = row.get("expected", "").strip().lower() == "duplicate"
upload_path = audio_path
tmp_resampled = None
try:
t0 = time.perf_counter()
......@@ -281,10 +319,10 @@ def main() -> None:
resampled = _resample_audio(audio_path, target_sr=args.sample_rate)
if resampled:
upload_path = resampled
tmp_files.append(resampled)
tmp_resampled = resampled
# 2. 上传到 OSS
query_id = f"query_{i}_{query_song_id}"
query_id = f"query_{idx}_{query_song_id}"
if not args.skip_upload:
oss_url = upload_to_oss(oss_bucket, upload_path, query_id)
else:
......@@ -294,12 +332,8 @@ def main() -> None:
# 3. 提交 DNA 查询
t1 = time.perf_counter()
if oss_url:
job_id = submit_dna_query(ice_client, oss_url, query_id)
else:
job_id = submit_dna_query(ice_client, "", query_id)
logger.info("[%d/%d] 提交查询: job_id=%s", i, len(rows), job_id)
job_id = submit_dna_query(ice_client, oss_url or "", query_id)
logger.info("[%d/%d] 提交查询: job_id=%s", idx, total_pending, job_id)
# 4. 轮询结果
result = poll_job_result(ice_client, job_id)
......@@ -320,9 +354,7 @@ def main() -> None:
pk = match.get("PrimaryKey", "")
sim = match.get("GlobalSimilarity", 0.0)
topk_song_ids.append((pk, sim))
topk_song_ids.sort(key=lambda x: x[1], reverse=True)
if topk_song_ids:
top1_song_id = topk_song_ids[0][0]
top1_sim = round(topk_song_ids[0][1], 4)
......@@ -342,7 +374,15 @@ def main() -> None:
correct = expected_dup == predicted_dup
result_rows.append({
logger.info(
"[%d/%d] variant=%s expected=%s predicted_dup=%s top1=%s sim=%s"
" top1_hit=%s topk_hit=%s correct=%s time=%dms",
idx, total_pending, row.get("variant", ""), row.get("expected", ""),
predicted_dup, top1_song_id or "-", top1_sim if top1_sim != "" else "-",
top1_hit, topk_hit, correct, total_ms,
)
return {
"query_song_id": query_song_id,
"audio_song_id": audio_song_id,
"audio_path": audio_path,
......@@ -363,19 +403,12 @@ def main() -> None:
"poll_ms": poll_ms,
"total_ms": total_ms,
"error": "",
})
logger.info(
"[%d/%d] variant=%s expected=%s predicted_dup=%s top1=%s sim=%s top1_hit=%s topk_hit=%s correct=%s time=%dms",
i, len(rows), row.get("variant", ""), row.get("expected", ""),
predicted_dup, top1_song_id or "-", top1_sim if top1_sim != "" else "-",
top1_hit, topk_hit, correct, total_ms,
)
}
except Exception as e:
total_ms = round((time.perf_counter() - t0) * 1000, 1)
logger.error("[%d/%d] 查询失败: %s, %s", i, len(rows), audio_path, e)
result_rows.append({
logger.error("[%d/%d] 查询失败: %s, %s", idx, total_pending, audio_path, e)
return {
"query_song_id": query_song_id,
"audio_song_id": audio_song_id,
"audio_path": audio_path,
......@@ -396,33 +429,22 @@ def main() -> None:
"poll_ms": "",
"total_ms": total_ms,
"error": str(e),
})
# 清理临时文件
for tmp in tmp_files:
try:
Path(tmp).unlink(missing_ok=True)
except Exception:
pass
# 写逐条结果
fieldnames = ["query_song_id", "audio_song_id", "audio_path", "variant", "sample_class",
"expected_song_id", "expected", "top1_song_id", "top1_similarity",
"top1_hit", "topk_hit", "expected_rank", "expected_similarity",
"expected_duplicate", "predicted_duplicate", "correct",
"upload_ms", "poll_ms", "total_ms", "error"]
}
finally:
if tmp_resampled:
Path(tmp_resampled).unlink(missing_ok=True)
with out_path.open("w", newline="", encoding="utf-8") as f:
writer = csv.DictWriter(f, fieldnames=fieldnames)
writer.writeheader()
writer.writerows(result_rows)
# 汇总辅助函数(需在 _collect/_flush_summary 调用前定义)
def _to_bool(v) -> bool:
if isinstance(v, bool):
return v
return str(v).strip().lower() in ("true", "1", "yes")
# 汇总指标
def _metrics(rows: list[dict]) -> dict:
tp = sum(1 for r in rows if r["expected_duplicate"] and r["predicted_duplicate"])
fp = sum(1 for r in rows if not r["expected_duplicate"] and r["predicted_duplicate"])
tn = sum(1 for r in rows if not r["expected_duplicate"] and not r["predicted_duplicate"])
fn = sum(1 for r in rows if r["expected_duplicate"] and not r["predicted_duplicate"])
tp = sum(1 for r in rows if _to_bool(r["expected_duplicate"]) and _to_bool(r["predicted_duplicate"]))
fp = sum(1 for r in rows if not _to_bool(r["expected_duplicate"]) and _to_bool(r["predicted_duplicate"]))
tn = sum(1 for r in rows if not _to_bool(r["expected_duplicate"]) and not _to_bool(r["predicted_duplicate"]))
fn = sum(1 for r in rows if _to_bool(r["expected_duplicate"]) and not _to_bool(r["predicted_duplicate"]))
precision = tp / (tp + fp) if tp + fp else 0.0
recall = tp / (tp + fn) if tp + fn else 0.0
f1 = 2 * precision * recall / (precision + recall) if precision + recall else 0.0
......@@ -436,21 +458,6 @@ def main() -> None:
"tp": tp, "fp": fp, "tn": tn, "fn": fn,
}
metrics = _metrics(result_rows)
from collections import defaultdict
by_variant: dict[str, dict] = defaultdict(lambda: {"correct": 0, "total": 0})
for r in result_rows:
v = r["variant"] or "unknown"
by_variant[v]["total"] += 1
if r["correct"]:
by_variant[v]["correct"] += 1
# 耗时统计
total_times = [r["total_ms"] for r in result_rows if r.get("total_ms", "") != ""]
upload_times = [r["upload_ms"] for r in result_rows if r.get("upload_ms", "") != ""]
poll_times = [r["poll_ms"] for r in result_rows if r.get("poll_ms", "") != ""]
def _time_stats(values: list) -> dict:
s = sorted(values) if values else []
n = len(s)
......@@ -463,33 +470,70 @@ def main() -> None:
"max": round(max(values), 1) if n else 0,
}
summary = {
"total": len(result_rows),
"filters": {
"variants": sorted(variant_filter) if variant_filter else None,
"sample_classes": sorted(sample_class_filter) if sample_class_filter else None,
"expected": args.expected,
"original_total": original_count,
},
"duplicate_threshold": args.duplicate_threshold,
"accuracy": metrics["accuracy"],
"precision": metrics["precision"],
"recall": metrics["recall"],
"f1": metrics["f1"],
"tp": metrics["tp"], "fp": metrics["fp"], "tn": metrics["tn"], "fn": metrics["fn"],
"query_time_ms": _time_stats(total_times),
"upload_time_ms": _time_stats(upload_times),
"poll_time_ms": _time_stats(poll_times),
"by_variant": {
v: {"accuracy": round(d["correct"] / d["total"], 4), "total": d["total"]}
for v, d in sorted(by_variant.items())
},
"out": str(out_path),
}
from collections import defaultdict
def _flush_summary(rows: list[dict]) -> None:
metrics = _metrics(rows)
by_variant: dict[str, dict] = defaultdict(lambda: {"correct": 0, "total": 0})
for r in rows:
v = r["variant"] or "unknown"
by_variant[v]["total"] += 1
if _to_bool(r["correct"]):
by_variant[v]["correct"] += 1
total_times = [float(r["total_ms"]) for r in rows if r.get("total_ms", "") != ""]
upload_times = [float(r["upload_ms"]) for r in rows if r.get("upload_ms", "") != ""]
poll_times = [float(r["poll_ms"]) for r in rows if r.get("poll_ms", "") != ""]
summary = {
"total": len(rows),
"pending": total_pending - len(rows),
"filters": {
"variants": sorted(variant_filter) if variant_filter else None,
"sample_classes": sorted(sample_class_filter) if sample_class_filter else None,
"expected": args.expected,
"original_total": original_count,
},
"duplicate_threshold": args.duplicate_threshold,
"accuracy": metrics["accuracy"],
"precision": metrics["precision"],
"recall": metrics["recall"],
"f1": metrics["f1"],
"tp": metrics["tp"], "fp": metrics["fp"], "tn": metrics["tn"], "fn": metrics["fn"],
"query_time_ms": _time_stats(total_times),
"upload_time_ms": _time_stats(upload_times),
"poll_time_ms": _time_stats(poll_times),
"by_variant": {
v: {"accuracy": round(d["correct"] / d["total"], 4), "total": d["total"]}
for v, d in sorted(by_variant.items())
},
"out": str(out_path),
}
summary_path = out_path.with_suffix(".summary.json")
summary_path.write_text(json.dumps(summary, ensure_ascii=False, indent=2), encoding="utf-8")
def _collect(result_dict: dict) -> None:
nonlocal completed_count
with lock:
result_rows.append(result_dict)
completed_count += 1
if args.save_interval > 0 and completed_count % args.save_interval == 0:
_flush_results(result_rows)
_flush_summary(result_rows)
logger.info(" [checkpoint] 已保存 %d 条结果到 %s", len(result_rows), out_path)
with ThreadPoolExecutor(max_workers=args.concurrency) as executor:
futures = {
executor.submit(_process_one, row, idx): idx
for idx, row in enumerate(rows, 1)
}
for future in as_completed(futures):
_collect(future.result())
# 写逐条结果(最终落盘)
_flush_results(result_rows)
summary_path = out_path.with_suffix(".summary.json")
summary_path.write_text(json.dumps(summary, ensure_ascii=False, indent=2), encoding="utf-8")
print(json.dumps(summary, ensure_ascii=False, indent=2))
_flush_summary(result_rows)
print(json.dumps(json.loads(summary_path.read_text(encoding="utf-8")), ensure_ascii=False, indent=2))
if __name__ == "__main__":
......
#!/usr/bin/env python3
import argparse
import csv
import os
import re
import sys
from pathlib import Path
from urllib.parse import urlparse
from dotenv import load_dotenv
from qcloud_cos import CosConfig, CosS3Client
load_dotenv(Path(__file__).resolve().parents[1] / '.env')
AUDIO_EXTS = {'.mp3', '.wav', '.flac', '.m4a', '.aac', '.ogg', '.wma', '.ape', '.alac'}
def normalize_key(raw_url: str) -> str:
if not raw_url:
return ''
raw_url = raw_url.strip()
if re.match(r'^https?://', raw_url, re.I):
return urlparse(raw_url).path.lstrip('/')
return raw_url.lstrip('/')
def ext_ok(path: str) -> bool:
return Path(path).suffix.lower() in AUDIO_EXTS
def main():
ap = argparse.ArgumentParser()
ap.add_argument('--manifest', default='output_selection_budgeted/selected_files.csv')
ap.add_argument('--output-dir', default='downloads')
ap.add_argument('--types', default='1,7,8,11,16')
ap.add_argument('--song-limit', type=int, default=3)
ap.add_argument('--start-song-offset', type=int, default=0)
ap.add_argument('--overwrite', action='store_true')
ap.add_argument('--fail-log', default='download_failures.csv')
args = ap.parse_args()
region = os.environ['COS_REGION']
secret_id = os.environ['COS_SECRET_ID']
secret_key = os.environ['COS_SECRET_KEY']
bucket = os.environ['COS_BUCKET']
config = CosConfig(Region=region, SecretId=secret_id, SecretKey=secret_key)
client = CosS3Client(config)
wanted_types = {int(x) for x in args.types.split(',') if x.strip()}
out_dir = Path(args.output_dir)
out_dir.mkdir(parents=True, exist_ok=True)
fail_log = Path(args.fail_log)
fail_exists = fail_log.exists()
fail_f = fail_log.open('a', encoding='utf-8', newline='')
fail_w = csv.writer(fail_f)
if not fail_exists:
fail_w.writerow(['song_id', 'type', 'url', 'error'])
downloaded = skipped = errors = 0
active_song_ids = []
current_song = None
started = False
song_counter = 0
with open(args.manifest, 'r', encoding='utf-8', newline='') as f:
reader = csv.DictReader(f)
for row in reader:
sid = row['song_id']
if sid != current_song:
current_song = sid
if song_counter < args.start_song_offset:
song_counter += 1
continue
if args.song_limit and started and len(active_song_ids) >= args.song_limit and sid not in set(active_song_ids):
break
if sid not in active_song_ids:
active_song_ids.append(sid)
started = True
if sid not in active_song_ids:
continue
try:
t = int(row['type'])
except Exception:
skipped += 1
continue
if t not in wanted_types:
skipped += 1
continue
key = normalize_key(row['url'])
if not key or not ext_ok(key):
skipped += 1
continue
target_dir = out_dir / sid / f'type_{t}'
target_dir.mkdir(parents=True, exist_ok=True)
target = target_dir / Path(key).name
if target.exists() and not args.overwrite:
skipped += 1
continue
try:
client.download_file(Bucket=bucket, Key=key, DestFilePath=str(target))
downloaded += 1
if downloaded % 50 == 0:
print({'downloaded': downloaded, 'songs': len(active_song_ids), 'last': str(target)})
except Exception as e:
errors += 1
fail_w.writerow([sid, t, row['url'], str(e)])
fail_f.flush()
print({'error_song': sid, 'type': t, 'key': key, 'error': str(e)}, file=sys.stderr)
fail_f.close()
print({
'downloaded_files': downloaded,
'skipped_rows': skipped,
'error_files': errors,
'song_count': len(active_song_ids),
'output_dir': str(out_dir.resolve()),
})
if __name__ == '__main__':
main()
#!/usr/bin/env python3
"""从 embed_db 数据库下载歌曲音频,保留完整元数据。
用法:
# 只下载主版本,限 100 首歌
python scripts/download_from_db.py --song-limit 100
# 主版本 + 每首歌最多 3 个其他版本
python scripts/download_from_db.py --song-limit 100 --extra-versions 3
# 主版本 + 所有其他版本
python scripts/download_from_db.py --song-limit 100 --extra-versions -1
# 指定歌曲 ID
python scripts/download_from_db.py --song-ids 1,2,3 --extra-versions -1
# 本地:导出查询结果到 CSV(不下载)
python scripts/download_from_db.py --song-limit 500 --extra-versions 3 --export-csv records.csv
# 服务器:从导出的 CSV 下载(不需要连接数据库)
python scripts/download_from_db.py --from-csv records.csv --concurrency 8
输出目录结构:
{output_dir}/
metadata.csv — 所有下载记录的完整元数据
reference.csv — DNA 入库参照集(is_main_version=1 的主版本)
queries.csv — DNA 评测查询集(其余版本,expected=duplicate)
download_failures.csv — 下载失败记录
{song_id}_{safe_song_name}/
{record_id}_{safe_singer_name}{ext}
"""
import argparse
import csv
import logging
import re
import time
import threading
import urllib.parse
import urllib.request
from concurrent.futures import ThreadPoolExecutor, as_completed
from pathlib import Path
from dotenv import load_dotenv
try:
from tqdm import tqdm
except ImportError:
tqdm = None
load_dotenv(Path(__file__).resolve().parent.parent / ".env")
logger = logging.getLogger(__name__)
DB_DSN = "postgresql://postgres:postgres@localhost:5432/embed_db"
# fetch_records 查询返回的列,也是 --export-csv / --from-csv 的 CSV 格式
RECORDS_FIELDS = [
"record_id", "song_id", "song_name", "original_singer",
"lyricist", "composer", "genre", "language_tag",
"singer_name", "version_name", "platform_name",
"duration", "is_main_version", "album", "pub_time", "audio_url",
]
METADATA_FIELDS = [
"record_id", "song_id", "song_name", "original_singer",
"lyricist", "composer", "genre", "language_tag",
"singer_name", "version_name", "platform_name",
"duration", "is_main_version", "album", "pub_time",
"audio_url", "local_path", "status",
]
REFERENCE_FIELDS = [
"song_id", "audio_path", "variant",
"song_name", "original_singer", "lyricist", "composer",
"genre", "language_tag",
]
QUERY_FIELDS = [
"song_id", "audio_path", "variant", "sample_class",
"expected_song_id", "expected",
"song_name", "original_singer", "singer_name",
"platform_name", "duration", "is_main_version",
]
def _safe_name(s: str, max_len: int = 40) -> str:
if not s:
return "unknown"
s = re.sub(r'[\\/:*?"<>|]', "_", s).strip()
return s[:max_len] or "unknown"
def _ext_from_url(url: str) -> str:
path = url.split("?")[0]
suffix = Path(path).suffix.lower()
return suffix if suffix in {".mp3", ".wav", ".flac", ".m4a", ".aac", ".ogg"} else ".mp3"
def _encode_url(url: str) -> str:
parsed = urllib.parse.urlsplit(url)
encoded_path = urllib.parse.quote(parsed.path, safe="/:@!$&'()*+,;=")
return urllib.parse.urlunsplit(parsed._replace(path=encoded_path))
def fetch_records(
conn: psycopg.Connection,
song_ids: list[int] | None,
extra_versions: int,
song_limit: int,
song_offset: int,
) -> list[dict]:
"""查询待下载记录,JOIN embed_song 获取完整元数据。
extra_versions:
0 — 只下载主版本(is_main_version=1)
-1 — 主版本 + 所有其他版本
N — 主版本 + 每首歌最多 N 个其他版本(按 id 排序取前 N)
"""
base_where = "r.audio_url IS NOT NULL AND r.audio_url != '' AND r.status = 'ready'"
if song_ids:
ids_str = ",".join(str(i) for i in song_ids)
base_where += f" AND r.song_id IN ({ids_str})"
# 确定要下载的 song_id 范围
if song_limit > 0 and not song_ids:
song_range_sql = f"""
SELECT DISTINCT song_id FROM embed_record
WHERE {base_where}
ORDER BY song_id
LIMIT {song_limit} OFFSET {song_offset}
"""
song_range_clause = f"r.song_id IN ({song_range_sql})"
else:
song_range_clause = "TRUE"
full_where = f"{base_where} AND {song_range_clause}"
# 主版本
main_sql = f"""
SELECT
r.id AS record_id, r.song_id,
s.song_name, s.original_singer, s.lyricist, s.composer,
s.genre, s.language_tag,
r.singer_name, r.version_name, r.platform_name,
r.duration, r.is_main_version, r.album, r.pub_time, r.audio_url
FROM embed_record r
LEFT JOIN embed_song s ON s.id = r.song_id
WHERE {full_where} AND r.is_main_version = 1
"""
if extra_versions == 0:
union_sql = main_sql
elif extra_versions == -1:
extra_sql = f"""
SELECT
r.id AS record_id, r.song_id,
s.song_name, s.original_singer, s.lyricist, s.composer,
s.genre, s.language_tag,
r.singer_name, r.version_name, r.platform_name,
r.duration, r.is_main_version, r.album, r.pub_time, r.audio_url
FROM embed_record r
LEFT JOIN embed_song s ON s.id = r.song_id
WHERE {full_where} AND r.is_main_version = 0
"""
union_sql = f"{main_sql} UNION ALL {extra_sql}"
else:
# 每首歌最多取 extra_versions 个非主版本,用窗口函数在 SQL 层截断
extra_sql = f"""
SELECT record_id, song_id, song_name, original_singer, lyricist, composer,
genre, language_tag, singer_name, version_name, platform_name,
duration, is_main_version, album, pub_time, audio_url
FROM (
SELECT
r.id AS record_id, r.song_id,
s.song_name, s.original_singer, s.lyricist, s.composer,
s.genre, s.language_tag,
r.singer_name, r.version_name, r.platform_name,
r.duration, r.is_main_version, r.album, r.pub_time, r.audio_url,
ROW_NUMBER() OVER (PARTITION BY r.song_id ORDER BY r.id) AS rn
FROM embed_record r
LEFT JOIN embed_song s ON s.id = r.song_id
WHERE {full_where} AND r.is_main_version = 0
) ranked
WHERE rn <= {extra_versions}
"""
union_sql = f"{main_sql} UNION ALL {extra_sql}"
final_sql = f"""
SELECT * FROM ({union_sql}) combined
ORDER BY song_id, is_main_version DESC, record_id
"""
with conn.cursor() as cur:
cur.execute(final_sql)
cols = [desc[0] for desc in cur.description]
return [dict(zip(cols, row)) for row in cur.fetchall()]
def load_records_from_csv(path: str) -> list[dict]:
"""从 --export-csv 导出的文件读取记录,替代数据库查询。"""
with open(path, newline="", encoding="utf-8") as f:
rows = list(csv.DictReader(f))
# 将字符串还原为适当类型
for r in rows:
r["record_id"] = int(r["record_id"])
r["song_id"] = int(r["song_id"])
r["is_main_version"] = int(r["is_main_version"]) if r.get("is_main_version") else 0
r["duration"] = int(r["duration"]) if r.get("duration") else None
return rows
def download_audio(url: str, dest: Path, timeout: int = 60, retries: int = 2) -> bool:
url = _encode_url(url)
for attempt in range(retries + 1):
try:
req = urllib.request.Request(url, headers={"User-Agent": "Mozilla/5.0"})
with urllib.request.urlopen(req, timeout=timeout) as resp:
data = resp.read()
dest.write_bytes(data)
return True
except Exception as e:
if attempt < retries:
time.sleep(2 ** attempt)
else:
logger.warning("下载失败 [%d/%d]: %s — %s", attempt + 1, retries + 1, url, e)
return False
def _build_paths(record: dict, output_dir: Path) -> tuple[Path, Path]:
song_dir = output_dir / f"{record['song_id']}_{_safe_name(record['song_name'] or '')}"
ext = _ext_from_url(record["audio_url"])
singer = _safe_name(record["singer_name"] or record["original_singer"] or "unknown")
filename = f"{record['record_id']}_{singer}{ext}"
return song_dir, song_dir / filename
def main() -> None:
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s [%(levelname)s] %(message)s",
)
ap = argparse.ArgumentParser(description="从 embed_db 下载歌曲音频并生成测试集 CSV")
ap.add_argument("--output-dir", default="downloads_db", help="输出根目录(默认 downloads_db)")
ap.add_argument("--song-limit", type=int, default=0, help="限制歌曲数量(0=不限)")
ap.add_argument("--song-offset", type=int, default=0, help="按 song_id 跳过前 N 首歌")
ap.add_argument("--song-ids", help="只下载指定 song_id,逗号分隔")
ap.add_argument(
"--extra-versions", type=int, default=0, metavar="N",
help="每首歌额外下载 N 个非主版本(0=只主版本,-1=全部,默认 0)",
)
ap.add_argument("--concurrency", type=int, default=4, help="并发下载线程数(默认 4)")
ap.add_argument("--timeout", type=int, default=60, help="单文件下载超时秒数(默认 60)")
ap.add_argument("--overwrite", action="store_true", help="覆盖已存在的文件")
ap.add_argument("--dsn", default=DB_DSN, help="PostgreSQL 连接串")
ap.add_argument(
"--export-csv", metavar="FILE",
help="只导出查询结果到 CSV 后退出,不执行下载(在本地机器上运行)",
)
ap.add_argument(
"--from-csv", metavar="FILE",
help="从已导出的 CSV 读取记录,跳过数据库连接(在服务器上运行)",
)
args = ap.parse_args()
song_ids = [int(x) for x in args.song_ids.split(",")] if args.song_ids else None
# ── 模式一:从数据库查询 ──────────────────────────────────────────────────
if args.from_csv:
logger.info("从 CSV 读取记录: %s", args.from_csv)
records = load_records_from_csv(args.from_csv)
else:
logger.info("连接数据库: %s", args.dsn)
try:
import psycopg
except ImportError:
logger.error("缺少 psycopg,请执行: pip install psycopg")
raise SystemExit(1)
with psycopg.connect(args.dsn) as conn:
records = fetch_records(
conn,
song_ids=song_ids,
extra_versions=args.extra_versions,
song_limit=args.song_limit,
song_offset=args.song_offset,
)
song_count_total = len({r["song_id"] for r in records})
logger.info("共 %d 条记录,涉及 %d 首歌", len(records), song_count_total)
# ── 模式二:只导出 CSV,不下载 ────────────────────────────────────────────
if args.export_csv:
out_csv = Path(args.export_csv)
out_csv.parent.mkdir(parents=True, exist_ok=True)
with out_csv.open("w", newline="", encoding="utf-8") as f:
writer = csv.DictWriter(f, fieldnames=RECORDS_FIELDS, extrasaction="ignore")
writer.writeheader()
writer.writerows(records)
logger.info("已导出 %d 条记录到 %s,可同步到服务器后用 --from-csv 下载", len(records), out_csv)
return
# ── 模式三:下载 ──────────────────────────────────────────────────────────
output_dir = Path(args.output_dir)
output_dir.mkdir(parents=True, exist_ok=True)
# 断点续跑:加载已有 metadata.csv
metadata_path = output_dir / "metadata.csv"
done_record_ids: set[int] = set()
existing_metadata: list[dict] = []
if metadata_path.exists():
with metadata_path.open(newline="", encoding="utf-8") as f:
for row in csv.DictReader(f):
if row.get("status") == "ok":
done_record_ids.add(int(row["record_id"]))
existing_metadata.append(row)
logger.info("断点续跑:跳过已下载 %d 条", len(done_record_ids))
pending = [r for r in records if r["record_id"] not in done_record_ids]
logger.info("本次待下载: %d 条", len(pending))
for rec in pending:
song_dir, dest = _build_paths(rec, output_dir)
rec["_song_dir"] = song_dir
rec["_dest"] = dest
results: list[dict] = list(existing_metadata)
results_lock = threading.Lock()
def _download_one(rec: dict) -> dict:
song_dir: Path = rec["_song_dir"]
dest: Path = rec["_dest"]
song_dir.mkdir(parents=True, exist_ok=True)
if dest.exists() and not args.overwrite:
status = "ok"
else:
ok = download_audio(rec["audio_url"], dest, timeout=args.timeout)
status = "ok" if ok else "failed"
return {
"record_id": rec["record_id"],
"song_id": rec["song_id"],
"song_name": rec["song_name"] or "",
"original_singer": rec["original_singer"] or "",
"lyricist": rec["lyricist"] or "",
"composer": rec["composer"] or "",
"genre": rec["genre"] or "",
"language_tag": rec["language_tag"] or "",
"singer_name": rec["singer_name"] or "",
"version_name": rec["version_name"] or "",
"platform_name": rec["platform_name"] or "",
"duration": rec["duration"] or "",
"is_main_version": rec["is_main_version"],
"album": rec["album"] or "",
"pub_time": rec["pub_time"] or "",
"audio_url": rec["audio_url"],
"local_path": str(dest) if status == "ok" else "",
"status": status,
}
progress = (
tqdm(total=len(pending), unit="文件", dynamic_ncols=True)
if tqdm is not None
else None
)
with ThreadPoolExecutor(max_workers=args.concurrency) as executor:
futures = {executor.submit(_download_one, rec): rec for rec in pending}
for future in as_completed(futures):
result = future.result()
with results_lock:
results.append(result)
if progress is not None:
status_str = "✓" if result["status"] == "ok" else "✗"
progress.set_postfix_str(
f"{status_str} {result['song_name']} — {result['singer_name'] or result['original_singer']}",
refresh=False,
)
progress.update(1)
else:
done = sum(1 for r in results if r not in existing_metadata)
if done % 20 == 0 or done == len(pending):
logger.info("[%d/%d] %s — %s (%s)",
done, len(pending),
result["song_name"], result["singer_name"], result["status"])
if progress is not None:
progress.close()
# 写 metadata.csv
with metadata_path.open("w", newline="", encoding="utf-8") as f:
writer = csv.DictWriter(f, fieldnames=METADATA_FIELDS)
writer.writeheader()
writer.writerows(results)
# 写 download_failures.csv
failures = [r for r in results if r.get("status") == "failed"]
if failures:
fail_path = output_dir / "download_failures.csv"
with fail_path.open("w", newline="", encoding="utf-8") as f:
writer = csv.DictWriter(f, fieldnames=METADATA_FIELDS, extrasaction="ignore")
writer.writeheader()
writer.writerows(failures)
logger.warning("失败 %d 条,见 %s", len(failures), fail_path)
# 生成 reference.csv 和 queries.csv
ok_results = [r for r in results if r.get("status") == "ok" and r.get("local_path")]
ref_rows, query_rows = [], []
for r in ok_results:
if int(r["is_main_version"]) == 1:
ref_rows.append({
"song_id": r["song_id"],
"audio_path": r["local_path"],
"variant": "original",
"song_name": r["song_name"],
"original_singer": r["original_singer"],
"lyricist": r["lyricist"],
"composer": r["composer"],
"genre": r["genre"],
"language_tag": r["language_tag"],
})
else:
platform = r["platform_name"] or "unknown"
query_rows.append({
"song_id": r["song_id"],
"audio_path": r["local_path"],
"variant": f"cover_{platform}",
"sample_class": "positive",
"expected_song_id": r["song_id"],
"expected": "duplicate",
"song_name": r["song_name"],
"original_singer": r["original_singer"],
"singer_name": r["singer_name"],
"platform_name": r["platform_name"],
"duration": r["duration"],
"is_main_version": r["is_main_version"],
})
ref_path = output_dir / "reference.csv"
with ref_path.open("w", newline="", encoding="utf-8") as f:
writer = csv.DictWriter(f, fieldnames=REFERENCE_FIELDS)
writer.writeheader()
writer.writerows(ref_rows)
query_path = output_dir / "queries.csv"
with query_path.open("w", newline="", encoding="utf-8") as f:
writer = csv.DictWriter(f, fieldnames=QUERY_FIELDS)
writer.writeheader()
writer.writerows(query_rows)
ok_count = sum(1 for r in results if r.get("status") == "ok")
song_count_ok = len({r["song_id"] for r in results if r.get("status") == "ok"})
logger.info("完成: 成功 %d 条,失败 %d 条,涉及 %d 首歌",
ok_count, len(failures), song_count_ok)
logger.info("参照集: %s(%d 条)", ref_path, len(ref_rows))
logger.info("查询集: %s(%d 条)", query_path, len(query_rows))
logger.info("元数据: %s", metadata_path)
if __name__ == "__main__":
main()