Commit 64d5ce50 64d5ce50acd984d6de820d90f7b068ce4a69770f by 沈秋雨

refactor(deps): 修改导入路径以适配新包结构

- 将lyric_dedup_server相关导入改为dedup_server
- 统一调整多个脚本和测试文件中的导入路径
- 添加脚本aliyun_dna/test_single_pair.py,实现单对音频DNA匹配测试功能
- 新增上传、提交作业、轮询及结果解析等关键步骤的实现
- 支持通过命令行参数控制测试流程及阈值设置
- 增强日志记录和错误处理,提升测试可靠性
1 parent 03dd9219
......@@ -54,8 +54,8 @@ def main() -> None:
def check_file_pg(args: argparse.Namespace) -> None:
from lyric_dedup_server.config import ServerConfig
from lyric_dedup_server.service import DedupService
from dedup_server.config import ServerConfig
from dedup_server.service import DedupService
record = record_from_file(Path(args.file))
config = ServerConfig(
......
#!/usr/bin/env python3
"""单对音频 DNA 匹配测试:将原版入库,再以翻唱查询,输出相似度结果。
用法:
python scripts/aliyun_dna/test_single_pair.py \
--ref "/path/to/原版.mp3" --ref-id original_id \
--query "/path/to/翻唱.mp3" --query-id cover_id \
--db-id <DNA库ID> # 可选,默认读 .env
"""
import argparse
import json
import logging
import os
import sys
import time
import urllib.request
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parents[2]))
from dotenv import load_dotenv
load_dotenv(Path(__file__).resolve().parents[2] / ".env")
import oss2
from alibabacloud_ice20201109.client import Client as ICEClient
from alibabacloud_tea_openapi.models import Config as OpenApiConfig
from alibabacloud_ice20201109 import models as ice_models
logging.basicConfig(level=logging.INFO, format="%(asctime)s [%(levelname)s] %(message)s")
logger = logging.getLogger(__name__)
POLL_INTERVAL = 3
POLL_TIMEOUT = 300
def _get_oss_client() -> oss2.Bucket:
auth = oss2.Auth(os.environ["ALIYUN_ICE_ACCESS_KEY_ID"], os.environ["ALIYUN_ICE_ACCESS_KEY_SECRET"])
return oss2.Bucket(auth, f"https://{os.environ['ALIYUN_OSS_ENDPOINT']}", os.environ["ALIYUN_OSS_BUCKET"])
def _get_ice_client() -> ICEClient:
region = os.environ.get("ALIYUN_ICE_REGION", "cn-hangzhou")
config = OpenApiConfig(
access_key_id=os.environ["ALIYUN_ICE_ACCESS_KEY_ID"],
access_key_secret=os.environ["ALIYUN_ICE_ACCESS_KEY_SECRET"],
endpoint=f"ice.{region}.aliyuncs.com",
region_id=region,
)
return ICEClient(config)
def upload_to_oss(bucket: oss2.Bucket, audio_path: str, prefix: str, key: str) -> str:
ext = Path(audio_path).suffix
obj = f"{prefix}/{key}{ext}"
size = Path(audio_path).stat().st_size
logger.info("上传: %s (%.1fMB) -> %s", Path(audio_path).name, size / 1e6, obj)
bucket.put_object_from_file(obj, audio_path)
return f"oss://{bucket.bucket_name}/{obj}"
def submit_job(client: ICEClient, oss_url: str, primary_key: str, save_type: str) -> str:
db_id = os.environ["ALIYUN_DNA_DB_ID"]
req = ice_models.SubmitDNAJobRequest(
input=ice_models.SubmitDNAJobRequestInput(type="OSS", media=oss_url),
primary_key=primary_key,
dbid=db_id,
config=json.dumps({"SaveType": save_type, "MediaType": "audio"}),
)
resp = client.submit_dnajob(req)
return resp.body.job_id
def poll_job(client: ICEClient, job_id: str) -> dict:
t0 = time.time()
while time.time() - t0 < POLL_TIMEOUT:
resp = client.query_dnajob_list(ice_models.QueryDNAJobListRequest(job_ids=job_id))
if resp.body.job_list:
job = resp.body.job_list[0]
if job.status == "Success":
return {"status": "Success", "dna_result_url": job.dnaresult}
elif job.status == "Fail":
return {"status": "Fail", "error": job.message or "unknown"}
time.sleep(POLL_INTERVAL)
return {"status": "Timeout"}
def fetch_results(url: str) -> list[dict]:
with urllib.request.urlopen(urllib.request.Request(url), timeout=30) as r:
data = json.loads(r.read().decode())
return data if isinstance(data, list) else [data]
def main():
ap = argparse.ArgumentParser(description="单对音频 DNA 匹配测试")
ap.add_argument("--ref", required=True, help="原版音频路径")
ap.add_argument("--ref-id", required=True, help="原版歌曲 ID(入库主键)")
ap.add_argument("--query", required=True, help="翻唱音频路径")
ap.add_argument("--query-id", default="query_cover", help="查询 ID")
ap.add_argument("--db-id", default="", help="覆盖 ALIYUN_DNA_DB_ID")
ap.add_argument("--skip-upload-ref", action="store_true", help="原版已入库,跳过上传")
ap.add_argument("--threshold", type=float, default=0.8, help="duplicate 判定阈值(默认 0.8)")
args = ap.parse_args()
if args.db_id:
os.environ["ALIYUN_DNA_DB_ID"] = args.db_id
for k in ["ALIYUN_ICE_ACCESS_KEY_ID", "ALIYUN_ICE_ACCESS_KEY_SECRET",
"ALIYUN_OSS_BUCKET", "ALIYUN_OSS_ENDPOINT", "ALIYUN_DNA_DB_ID"]:
if not os.environ.get(k):
logger.error("缺少环境变量: %s", k)
sys.exit(1)
logger.info("DNA 库: %s", os.environ["ALIYUN_DNA_DB_ID"])
bucket = _get_oss_client()
client = _get_ice_client()
# Step 1: 原版入库
if not args.skip_upload_ref:
logger.info("=== Step 1: 原版入库 (%s) ===", args.ref_id)
ref_oss = upload_to_oss(bucket, args.ref, "dna-test-ref", args.ref_id)
ref_job_id = submit_job(client, ref_oss, args.ref_id, "save")
logger.info("入库作业已提交: job_id=%s,等待完成...", ref_job_id)
ref_result = poll_job(client, ref_job_id)
if ref_result["status"] != "Success":
logger.error("原版入库失败: %s", ref_result)
sys.exit(1)
logger.info("原版入库成功,等待 5s 让索引生效...")
time.sleep(5)
else:
logger.info("=== Step 1: 跳过原版入库(--skip-upload-ref)===")
# Step 2: 翻唱查询
logger.info("=== Step 2: 翻唱查询 (%s) ===", args.query_id)
query_oss = upload_to_oss(bucket, args.query, "dna-test-query", args.query_id)
query_job_id = submit_job(client, query_oss, args.query_id, "nosave")
logger.info("查询作业已提交: job_id=%s,等待完成...", query_job_id)
query_result = poll_job(client, query_job_id)
if query_result["status"] != "Success":
logger.error("查询失败: %s", query_result)
sys.exit(1)
# Step 3: 解析结果
matches = fetch_results(query_result["dna_result_url"]) if query_result.get("dna_result_url") else []
matches.sort(key=lambda x: x.get("GlobalSimilarity", 0), reverse=True)
print("\n" + "=" * 60)
print(f"查询: {Path(args.query).name}")
print(f"原版 ID: {args.ref_id} | DNA库: {os.environ['ALIYUN_DNA_DB_ID']}")
print("=" * 60)
if not matches:
print("无匹配结果")
else:
for i, m in enumerate(matches[:5], 1):
pk = m.get("PrimaryKey", "")
sim = m.get("GlobalSimilarity", 0)
is_target = pk == args.ref_id
mark = " <-- 目标原版" if is_target else ""
print(f" #{i} PrimaryKey={pk} GlobalSimilarity={sim:.4f}{mark}")
top1 = matches[0]
top1_sim = top1.get("GlobalSimilarity", 0)
top1_pk = top1.get("PrimaryKey", "")
hit = top1_pk == args.ref_id
is_dup = top1_sim >= args.threshold
print()
print(f"Top1 命中原版: {'是' if hit else '否'}")
print(f"Top1 相似度: {top1_sim:.4f}(阈值 {args.threshold})")
print(f"判定为重复: {'是 ✓' if is_dup else '否 ✗'}")
print("=" * 60)
if __name__ == "__main__":
main()
......@@ -22,7 +22,7 @@ from lyric_dedup.file_import import read_lyric_file
from lyric_dedup.file_import import record_from_file
from lyric_dedup.normalization import fingerprint_text
from lyric_dedup.normalization import normalize_lyrics
from lyric_dedup_server.config import ServerConfig
from dedup_server.config import ServerConfig
def main() -> None:
......
......@@ -7,8 +7,8 @@ from lyric_dedup import LyricRecord
from lyric_dedup.eval_dataset import generate_eval_set
from lyric_dedup.file_import import record_from_file
from lyric_dedup.normalization import normalize_lyrics
from lyric_dedup_server.config import ServerConfig
from lyric_dedup_server.service import DedupService
from dedup_server.config import ServerConfig
from dedup_server.service import DedupService
BASE_LYRIC = """
......