import_audio_composition.py
7.29 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
178
179
180
181
182
183
184
185
186
187
188
189
190
191
"""批量导入音频文件到 composition_feature 表。
用法:
python scripts/import_audio_composition.py \
--dsn "postgresql:///lyric_dedup" \
--audio-dir /Volumes/移动硬盘/composition_test \
--ext .wav
支持通过 --file-list 指定一个包含音频路径的文本文件(每行一个路径)。
--update-chroma-full 模式:
仅为 chroma_full 列为 NULL 的已入库歌曲补写固定帧率 Chromagram,
跳过 feature_vector 重提取和 Dejavu 指纹计算,速度更快。
适用于开启子序列 DTW(COMPOSITION_SUBSEQUENCE_DTW_ENABLED=true)后的
首次全量补全,无需清空重建。
"""
import argparse
import logging
import sys
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parent.parent.parent))
from dotenv import load_dotenv
load_dotenv(Path(__file__).resolve().parent.parent.parent / ".env")
from tqdm import tqdm
from composition_dedup.service import CompositionConfig, CompositionDedupService
logger = logging.getLogger(__name__)
SUPPORTED_EXTENSIONS = {".mp3", ".wav", ".flac", ".ogg", ".m4a", ".aac", ".wma"}
def discover_audio_files(audio_dir: str | None, file_list: str | None, ext: str) -> list[tuple[str, str]]:
"""发现音频文件,返回 [(song_id, audio_path), ...] 列表。
优先使用 --file-list,否则扫描 --audio-dir 目录。
song_id 使用文件名的数字部分或路径的哈希值。
"""
results = []
if file_list:
with open(file_list, "r", encoding="utf-8") as f:
for line in f:
path = line.strip()
if not path:
continue
song_id = _extract_song_id(path)
results.append((song_id, path))
elif audio_dir:
audio_dir_path = Path(audio_dir)
for audio_file in sorted(audio_dir_path.rglob(f"*{ext}")):
if audio_file.is_file() and not audio_file.name.startswith("._"):
song_id = _extract_song_id(str(audio_file))
results.append((song_id, str(audio_file)))
else:
print("错误: 请指定 --audio-dir 或 --file-list")
sys.exit(1)
return results
def _extract_song_id(path: str) -> str:
"""从路径中提取 song_id。
优先取文件名第一段(下划线前),若为纯数字则使用,否则用路径哈希。
"""
name = Path(path).stem
prefix = name.split("_")[0]
if prefix.isdigit():
return prefix
import hashlib
return str(int(hashlib.md5(path.encode()).hexdigest()[:8], 16))
def main() -> None:
logging.basicConfig(level=logging.INFO, format="%(asctime)s [%(levelname)s] %(message)s")
parser = argparse.ArgumentParser(description="批量导入音频文件到 composition_feature 表")
parser.add_argument("--dsn", required=True, help="PostgreSQL DSN 连接串")
parser.add_argument("--audio-dir", help="音频文件目录")
parser.add_argument("--file-list", help="音频文件路径列表文件")
parser.add_argument("--ext", default=".wav", help="音频文件扩展名(默认 .wav)")
parser.add_argument("--batch-size", type=int, default=10, help="批次大小(默认 10)")
parser.add_argument("--clear", action="store_true", help="导入前清空 composition_feature 和 dejavu_fingerprints 表数据(保留表结构)")
parser.add_argument(
"--update-chroma-full",
action="store_true",
help="仅为 chroma_full=NULL 的已入库歌曲补写固定帧率 Chromagram,"
"跳过 feature_vector 重提取和 Dejavu 指纹,速度更快。"
"需先执行 migrate_add_chroma_full.sql 并确认 COMPOSITION_SUBSEQUENCE_DTW_ENABLED=true。",
)
args = parser.parse_args()
config = CompositionConfig(dsn=args.dsn)
service = CompositionDedupService(config=config)
if args.clear:
import psycopg
with psycopg.connect(args.dsn) as conn:
with conn.cursor() as cur:
cur.execute("TRUNCATE TABLE composition_feature, dejavu_fingerprints")
conn.commit()
logger.info("已清空 composition_feature 和 dejavu_fingerprints 表")
audio_files = discover_audio_files(args.audio_dir, args.file_list, args.ext)
logger.info("发现 %d 个音频文件", len(audio_files))
if args.update_chroma_full:
_update_chroma_full(config, audio_files)
return
success_count = 0
fail_count = 0
for start in tqdm(range(0, len(audio_files), args.batch_size), desc="导入进度"):
batch = audio_files[start:start + args.batch_size]
for song_id, audio_path in batch:
try:
service.ingest(song_id=int(song_id), audio_path=audio_path)
success_count += 1
except Exception as e:
logger.error("导入失败: song_id=%s, path=%s, error=%s", song_id, audio_path, e)
fail_count += 1
logger.info("导入完成: 成功 %d, 失败 %d", success_count, fail_count)
def _update_chroma_full(config: "CompositionConfig", audio_files: list[tuple[str, str]]) -> None:
"""仅补写 chroma_full 列,跳过 feature_vector 和 Dejavu 计算。"""
import psycopg
from composition_dedup.extractor import (
TARGET_SR,
extract_chroma_fixed_fps_from_samples,
load_audio_mono_22050hz,
)
# 查询库中 chroma_full 为 NULL 的 song_id
with psycopg.connect(config.dsn) as conn:
with conn.cursor() as cur:
cur.execute("SELECT song_id FROM composition_feature WHERE chroma_full IS NULL")
missing_ids = {str(row[0]) for row in cur.fetchall()}
if not missing_ids:
logger.info("所有已入库歌曲均已有 chroma_full,无需补全")
return
# 过滤:只处理缺少 chroma_full 的歌曲
to_update = [(sid, path) for sid, path in audio_files if sid in missing_ids]
logger.info(
"库中 chroma_full=NULL 的歌曲: %d 首,匹配到本地音频: %d 首(未匹配: %d 首)",
len(missing_ids),
len(to_update),
len(missing_ids) - len(to_update),
)
success_count = 0
fail_count = 0
for song_id, audio_path in tqdm(to_update, desc="补全 chroma_full"):
try:
samples = load_audio_mono_22050hz(audio_path)
chroma_full = extract_chroma_fixed_fps_from_samples(
samples, TARGET_SR,
target_fps=config.chroma_full_fps,
hop_length=config.chroma_hop_length,
win_len_smooth=config.chroma_win_len_smooth,
)
with psycopg.connect(config.dsn) as conn:
with conn.cursor() as cur:
cur.execute(
"""
UPDATE composition_feature
SET chroma_full = %s, chroma_n_frames = %s
WHERE song_id = %s
""",
(chroma_full.flatten().tolist(), int(chroma_full.shape[1]), int(song_id)),
)
conn.commit()
success_count += 1
except Exception as e:
logger.error("补全失败: song_id=%s, path=%s, error=%s", song_id, audio_path, e)
fail_count += 1
logger.info("chroma_full 补全完成: 成功 %d, 失败 %d", success_count, fail_count)
if __name__ == "__main__":
main()