refactor(deps): 修改导入路径以适配新包结构
- 将lyric_dedup_server相关导入改为dedup_server - 统一调整多个脚本和测试文件中的导入路径 - 添加脚本aliyun_dna/test_single_pair.py,实现单对音频DNA匹配测试功能 - 新增上传、提交作业、轮询及结果解析等关键步骤的实现 - 支持通过命令行参数控制测试流程及阈值设置 - 增强日志记录和错误处理,提升测试可靠性
Showing
4 changed files
with
182 additions
and
5 deletions
| ... | @@ -54,8 +54,8 @@ def main() -> None: | ... | @@ -54,8 +54,8 @@ def main() -> None: |
| 54 | 54 | ||
| 55 | 55 | ||
| 56 | def check_file_pg(args: argparse.Namespace) -> None: | 56 | def check_file_pg(args: argparse.Namespace) -> None: |
| 57 | from lyric_dedup_server.config import ServerConfig | 57 | from dedup_server.config import ServerConfig |
| 58 | from lyric_dedup_server.service import DedupService | 58 | from dedup_server.service import DedupService |
| 59 | 59 | ||
| 60 | record = record_from_file(Path(args.file)) | 60 | record = record_from_file(Path(args.file)) |
| 61 | config = ServerConfig( | 61 | config = ServerConfig( | ... | ... |
scripts/aliyun_dna/test_single_pair.py
0 → 100644
| 1 | #!/usr/bin/env python3 | ||
| 2 | """单对音频 DNA 匹配测试:将原版入库,再以翻唱查询,输出相似度结果。 | ||
| 3 | |||
| 4 | 用法: | ||
| 5 | python scripts/aliyun_dna/test_single_pair.py \ | ||
| 6 | --ref "/path/to/原版.mp3" --ref-id original_id \ | ||
| 7 | --query "/path/to/翻唱.mp3" --query-id cover_id \ | ||
| 8 | --db-id <DNA库ID> # 可选,默认读 .env | ||
| 9 | """ | ||
| 10 | |||
| 11 | import argparse | ||
| 12 | import json | ||
| 13 | import logging | ||
| 14 | import os | ||
| 15 | import sys | ||
| 16 | import time | ||
| 17 | import urllib.request | ||
| 18 | from pathlib import Path | ||
| 19 | |||
| 20 | sys.path.insert(0, str(Path(__file__).resolve().parents[2])) | ||
| 21 | from dotenv import load_dotenv | ||
| 22 | load_dotenv(Path(__file__).resolve().parents[2] / ".env") | ||
| 23 | |||
| 24 | import oss2 | ||
| 25 | from alibabacloud_ice20201109.client import Client as ICEClient | ||
| 26 | from alibabacloud_tea_openapi.models import Config as OpenApiConfig | ||
| 27 | from alibabacloud_ice20201109 import models as ice_models | ||
| 28 | |||
| 29 | logging.basicConfig(level=logging.INFO, format="%(asctime)s [%(levelname)s] %(message)s") | ||
| 30 | logger = logging.getLogger(__name__) | ||
| 31 | |||
| 32 | POLL_INTERVAL = 3 | ||
| 33 | POLL_TIMEOUT = 300 | ||
| 34 | |||
| 35 | |||
| 36 | def _get_oss_client() -> oss2.Bucket: | ||
| 37 | auth = oss2.Auth(os.environ["ALIYUN_ICE_ACCESS_KEY_ID"], os.environ["ALIYUN_ICE_ACCESS_KEY_SECRET"]) | ||
| 38 | return oss2.Bucket(auth, f"https://{os.environ['ALIYUN_OSS_ENDPOINT']}", os.environ["ALIYUN_OSS_BUCKET"]) | ||
| 39 | |||
| 40 | |||
| 41 | def _get_ice_client() -> ICEClient: | ||
| 42 | region = os.environ.get("ALIYUN_ICE_REGION", "cn-hangzhou") | ||
| 43 | config = OpenApiConfig( | ||
| 44 | access_key_id=os.environ["ALIYUN_ICE_ACCESS_KEY_ID"], | ||
| 45 | access_key_secret=os.environ["ALIYUN_ICE_ACCESS_KEY_SECRET"], | ||
| 46 | endpoint=f"ice.{region}.aliyuncs.com", | ||
| 47 | region_id=region, | ||
| 48 | ) | ||
| 49 | return ICEClient(config) | ||
| 50 | |||
| 51 | |||
| 52 | def upload_to_oss(bucket: oss2.Bucket, audio_path: str, prefix: str, key: str) -> str: | ||
| 53 | ext = Path(audio_path).suffix | ||
| 54 | obj = f"{prefix}/{key}{ext}" | ||
| 55 | size = Path(audio_path).stat().st_size | ||
| 56 | logger.info("上传: %s (%.1fMB) -> %s", Path(audio_path).name, size / 1e6, obj) | ||
| 57 | bucket.put_object_from_file(obj, audio_path) | ||
| 58 | return f"oss://{bucket.bucket_name}/{obj}" | ||
| 59 | |||
| 60 | |||
| 61 | def submit_job(client: ICEClient, oss_url: str, primary_key: str, save_type: str) -> str: | ||
| 62 | db_id = os.environ["ALIYUN_DNA_DB_ID"] | ||
| 63 | req = ice_models.SubmitDNAJobRequest( | ||
| 64 | input=ice_models.SubmitDNAJobRequestInput(type="OSS", media=oss_url), | ||
| 65 | primary_key=primary_key, | ||
| 66 | dbid=db_id, | ||
| 67 | config=json.dumps({"SaveType": save_type, "MediaType": "audio"}), | ||
| 68 | ) | ||
| 69 | resp = client.submit_dnajob(req) | ||
| 70 | return resp.body.job_id | ||
| 71 | |||
| 72 | |||
| 73 | def poll_job(client: ICEClient, job_id: str) -> dict: | ||
| 74 | t0 = time.time() | ||
| 75 | while time.time() - t0 < POLL_TIMEOUT: | ||
| 76 | resp = client.query_dnajob_list(ice_models.QueryDNAJobListRequest(job_ids=job_id)) | ||
| 77 | if resp.body.job_list: | ||
| 78 | job = resp.body.job_list[0] | ||
| 79 | if job.status == "Success": | ||
| 80 | return {"status": "Success", "dna_result_url": job.dnaresult} | ||
| 81 | elif job.status == "Fail": | ||
| 82 | return {"status": "Fail", "error": job.message or "unknown"} | ||
| 83 | time.sleep(POLL_INTERVAL) | ||
| 84 | return {"status": "Timeout"} | ||
| 85 | |||
| 86 | |||
| 87 | def fetch_results(url: str) -> list[dict]: | ||
| 88 | with urllib.request.urlopen(urllib.request.Request(url), timeout=30) as r: | ||
| 89 | data = json.loads(r.read().decode()) | ||
| 90 | return data if isinstance(data, list) else [data] | ||
| 91 | |||
| 92 | |||
| 93 | def main(): | ||
| 94 | ap = argparse.ArgumentParser(description="单对音频 DNA 匹配测试") | ||
| 95 | ap.add_argument("--ref", required=True, help="原版音频路径") | ||
| 96 | ap.add_argument("--ref-id", required=True, help="原版歌曲 ID(入库主键)") | ||
| 97 | ap.add_argument("--query", required=True, help="翻唱音频路径") | ||
| 98 | ap.add_argument("--query-id", default="query_cover", help="查询 ID") | ||
| 99 | ap.add_argument("--db-id", default="", help="覆盖 ALIYUN_DNA_DB_ID") | ||
| 100 | ap.add_argument("--skip-upload-ref", action="store_true", help="原版已入库,跳过上传") | ||
| 101 | ap.add_argument("--threshold", type=float, default=0.8, help="duplicate 判定阈值(默认 0.8)") | ||
| 102 | args = ap.parse_args() | ||
| 103 | |||
| 104 | if args.db_id: | ||
| 105 | os.environ["ALIYUN_DNA_DB_ID"] = args.db_id | ||
| 106 | |||
| 107 | for k in ["ALIYUN_ICE_ACCESS_KEY_ID", "ALIYUN_ICE_ACCESS_KEY_SECRET", | ||
| 108 | "ALIYUN_OSS_BUCKET", "ALIYUN_OSS_ENDPOINT", "ALIYUN_DNA_DB_ID"]: | ||
| 109 | if not os.environ.get(k): | ||
| 110 | logger.error("缺少环境变量: %s", k) | ||
| 111 | sys.exit(1) | ||
| 112 | |||
| 113 | logger.info("DNA 库: %s", os.environ["ALIYUN_DNA_DB_ID"]) | ||
| 114 | bucket = _get_oss_client() | ||
| 115 | client = _get_ice_client() | ||
| 116 | |||
| 117 | # Step 1: 原版入库 | ||
| 118 | if not args.skip_upload_ref: | ||
| 119 | logger.info("=== Step 1: 原版入库 (%s) ===", args.ref_id) | ||
| 120 | ref_oss = upload_to_oss(bucket, args.ref, "dna-test-ref", args.ref_id) | ||
| 121 | ref_job_id = submit_job(client, ref_oss, args.ref_id, "save") | ||
| 122 | logger.info("入库作业已提交: job_id=%s,等待完成...", ref_job_id) | ||
| 123 | ref_result = poll_job(client, ref_job_id) | ||
| 124 | if ref_result["status"] != "Success": | ||
| 125 | logger.error("原版入库失败: %s", ref_result) | ||
| 126 | sys.exit(1) | ||
| 127 | logger.info("原版入库成功,等待 5s 让索引生效...") | ||
| 128 | time.sleep(5) | ||
| 129 | else: | ||
| 130 | logger.info("=== Step 1: 跳过原版入库(--skip-upload-ref)===") | ||
| 131 | |||
| 132 | # Step 2: 翻唱查询 | ||
| 133 | logger.info("=== Step 2: 翻唱查询 (%s) ===", args.query_id) | ||
| 134 | query_oss = upload_to_oss(bucket, args.query, "dna-test-query", args.query_id) | ||
| 135 | query_job_id = submit_job(client, query_oss, args.query_id, "nosave") | ||
| 136 | logger.info("查询作业已提交: job_id=%s,等待完成...", query_job_id) | ||
| 137 | query_result = poll_job(client, query_job_id) | ||
| 138 | |||
| 139 | if query_result["status"] != "Success": | ||
| 140 | logger.error("查询失败: %s", query_result) | ||
| 141 | sys.exit(1) | ||
| 142 | |||
| 143 | # Step 3: 解析结果 | ||
| 144 | matches = fetch_results(query_result["dna_result_url"]) if query_result.get("dna_result_url") else [] | ||
| 145 | matches.sort(key=lambda x: x.get("GlobalSimilarity", 0), reverse=True) | ||
| 146 | |||
| 147 | print("\n" + "=" * 60) | ||
| 148 | print(f"查询: {Path(args.query).name}") | ||
| 149 | print(f"原版 ID: {args.ref_id} | DNA库: {os.environ['ALIYUN_DNA_DB_ID']}") | ||
| 150 | print("=" * 60) | ||
| 151 | |||
| 152 | if not matches: | ||
| 153 | print("无匹配结果") | ||
| 154 | else: | ||
| 155 | for i, m in enumerate(matches[:5], 1): | ||
| 156 | pk = m.get("PrimaryKey", "") | ||
| 157 | sim = m.get("GlobalSimilarity", 0) | ||
| 158 | is_target = pk == args.ref_id | ||
| 159 | mark = " <-- 目标原版" if is_target else "" | ||
| 160 | print(f" #{i} PrimaryKey={pk} GlobalSimilarity={sim:.4f}{mark}") | ||
| 161 | |||
| 162 | top1 = matches[0] | ||
| 163 | top1_sim = top1.get("GlobalSimilarity", 0) | ||
| 164 | top1_pk = top1.get("PrimaryKey", "") | ||
| 165 | hit = top1_pk == args.ref_id | ||
| 166 | is_dup = top1_sim >= args.threshold | ||
| 167 | |||
| 168 | print() | ||
| 169 | print(f"Top1 命中原版: {'是' if hit else '否'}") | ||
| 170 | print(f"Top1 相似度: {top1_sim:.4f}(阈值 {args.threshold})") | ||
| 171 | print(f"判定为重复: {'是 ✓' if is_dup else '否 ✗'}") | ||
| 172 | |||
| 173 | print("=" * 60) | ||
| 174 | |||
| 175 | |||
| 176 | if __name__ == "__main__": | ||
| 177 | main() |
| ... | @@ -22,7 +22,7 @@ from lyric_dedup.file_import import read_lyric_file | ... | @@ -22,7 +22,7 @@ from lyric_dedup.file_import import read_lyric_file |
| 22 | from lyric_dedup.file_import import record_from_file | 22 | from lyric_dedup.file_import import record_from_file |
| 23 | from lyric_dedup.normalization import fingerprint_text | 23 | from lyric_dedup.normalization import fingerprint_text |
| 24 | from lyric_dedup.normalization import normalize_lyrics | 24 | from lyric_dedup.normalization import normalize_lyrics |
| 25 | from lyric_dedup_server.config import ServerConfig | 25 | from dedup_server.config import ServerConfig |
| 26 | 26 | ||
| 27 | 27 | ||
| 28 | def main() -> None: | 28 | def main() -> None: | ... | ... |
| ... | @@ -7,8 +7,8 @@ from lyric_dedup import LyricRecord | ... | @@ -7,8 +7,8 @@ from lyric_dedup import LyricRecord |
| 7 | from lyric_dedup.eval_dataset import generate_eval_set | 7 | from lyric_dedup.eval_dataset import generate_eval_set |
| 8 | from lyric_dedup.file_import import record_from_file | 8 | from lyric_dedup.file_import import record_from_file |
| 9 | from lyric_dedup.normalization import normalize_lyrics | 9 | from lyric_dedup.normalization import normalize_lyrics |
| 10 | from lyric_dedup_server.config import ServerConfig | 10 | from dedup_server.config import ServerConfig |
| 11 | from lyric_dedup_server.service import DedupService | 11 | from dedup_server.service import DedupService |
| 12 | 12 | ||
| 13 | 13 | ||
| 14 | BASE_LYRIC = """ | 14 | BASE_LYRIC = """ | ... | ... |
-
Please register or sign in to post a comment