Skip to content
Toggle navigation
Toggle navigation
This project
Loading...
Sign in
沈秋雨
/
lyric_rhyme
Go to a project
Toggle navigation
Toggle navigation pinning
Projects
Groups
Snippets
Help
Project
Activity
Repository
Pipelines
Graphs
Issues
0
Merge Requests
0
Wiki
Network
Create a new issue
Builds
Commits
Issue Boards
Files
Commits
Network
Compare
Branches
Tags
Commit
63589d3c
...
63589d3c84acde4f51e11b0857525c12c41cbca2
authored
2026-06-27 10:28:06 +0800
by
沈秋雨
Browse Files
Options
Browse Files
Tag
Download
Email Patches
Plain Diff
更新支持pg下载多版本歌曲
1 parent
5ae39abf
Show whitespace changes
Inline
Side-by-side
Showing
4 changed files
with
787 additions
and
110 deletions
scripts/acrcloud/generate_acrcloud_testset.py
scripts/aliyun_dna/evaluate_aliyun_dna.py
scripts/download_cos.py
scripts/download_from_db.py
scripts/acrcloud/generate_acrcloud_testset.py
View file @
63589d3
...
...
@@ -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,20 +296,47 @@ 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
)
if
song_id
not
in
done_ref_ids
:
ref_rows
.
append
({
"song_id"
:
song_id
,
"audio_path"
:
str
(
wav
.
resolve
()),
"variant"
:
"original"
,
})
# 参照歌原始(自身查询,验证入库和识别链路)
audio_path_str
=
str
(
wav
.
resolve
())
if
audio_path_str
not
in
done_query_paths
:
query_rows
.
append
({
"song_id"
:
song_id
,
"audio_path"
:
str
(
wav
.
resolve
())
,
"audio_path"
:
audio_path_str
,
"variant"
:
"acr_original"
,
"sample_class"
:
"positive"
,
"expected_song_id"
:
song_id
,
...
...
@@ -317,8 +345,27 @@ def main() -> None:
for
variant_name
,
af
in
POSITIVE_VARIANTS
:
dst
=
variants_dir
/
f
"{song_id}_{variant_name}.wav"
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
:
...
...
@@ -326,7 +373,7 @@ def main() -> None:
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
:
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"
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"
])
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
)
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
(
"
%-20
s
%
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
(
"
%-25
s
%
d 条"
,
variant
,
cnt
)
if
__name__
==
"__main__"
:
main
()
...
...
scripts/aliyun_dna/evaluate_aliyun_dna.py
View file @
63589d3
...
...
@@ -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
))
out_path
=
Path
(
args
.
out
)
out_path
.
parent
.
mkdir
(
parents
=
True
,
exist_ok
=
True
)
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"
]
# 断点续跑:加载已处理的结果
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
()
out_path
=
Path
(
args
.
out
)
out_path
.
parent
.
mkdir
(
parents
=
True
,
exist_ok
=
True
)
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
)
tmp_files
=
[]
# 跟踪临时降采样文件
total_pending
=
len
(
rows
)
lock
=
threading
.
Lock
()
completed_count
=
0
result_rows
=
[]
for
i
,
row
in
enumerate
(
rows
,
1
):
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_{i
dx
}_{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
)
re
sult_rows
.
append
(
{
logger
.
error
(
"[
%
d/
%
d] 查询失败:
%
s,
%
s"
,
i
dx
,
total_pending
,
audio_path
,
e
)
re
turn
{
"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
),
})
}
finally
:
if
tmp_resampled
:
Path
(
tmp_resampled
)
.
unlink
(
missing_ok
=
True
)
# 清理临时文件
for
tmp
in
tmp_files
:
try
:
Path
(
tmp
)
.
unlink
(
missing_ok
=
True
)
except
Exception
:
pass
# 汇总辅助函数(需在 _collect/_flush_summary 调用前定义)
def
_to_bool
(
v
)
->
bool
:
if
isinstance
(
v
,
bool
):
return
v
return
str
(
v
)
.
strip
()
.
lower
()
in
(
"true"
,
"1"
,
"yes"
)
# 写逐条结果
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"
]
with
out_path
.
open
(
"w"
,
newline
=
""
,
encoding
=
"utf-8"
)
as
f
:
writer
=
csv
.
DictWriter
(
f
,
fieldnames
=
fieldnames
)
writer
.
writeheader
()
writer
.
writerows
(
result_rows
)
# 汇总指标
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,8 +470,22 @@ def main() -> None:
"max"
:
round
(
max
(
values
),
1
)
if
n
else
0
,
}
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
(
result_rows
),
"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
,
...
...
@@ -486,10 +507,33 @@ def main() -> None:
},
"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"
)
print
(
json
.
dumps
(
summary
,
ensure_ascii
=
False
,
indent
=
2
))
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"
)
_flush_summary
(
result_rows
)
print
(
json
.
dumps
(
json
.
loads
(
summary_path
.
read_text
(
encoding
=
"utf-8"
)),
ensure_ascii
=
False
,
indent
=
2
))
if
__name__
==
"__main__"
:
...
...
scripts/download_cos.py
0 → 100755
View file @
63589d3
#!/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
()
scripts/download_from_db.py
0 → 100755
View file @
63589d3
#!/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
()
Please
register
or
sign in
to post a comment