test_single_pair.py
6.79 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
#!/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()