#!/usr/bin/env python3
"""find_song_in_vcd.py —— 在 VCD 各段里滑动搜索视频中那首歌的音频位置

方法：把音频转成"每秒 12 维半音向量"（chroma），用视频中段的 chroma 序列
在每段 VCD 上滑动做归一化互相关，找峰值位置。对调性/音色/编曲差异鲁棒。

用法: python3 find_song_in_vcd.py --video V.mp4 --vstart 60 --vlen 30 --seg-dir MPEG目录
"""
import argparse, glob, math, os, struct, subprocess, cmath

ap = argparse.ArgumentParser()
ap.add_argument("--video", required=True)
ap.add_argument("--vstart", type=float, default=60.0, help="视频里取样的起点(秒)")
ap.add_argument("--vlen", type=float, default=30.0, help="取样长度(秒)")
ap.add_argument("--seg-dir", required=True)
ap.add_argument("--sr", type=int, default=11025)
ap.add_argument("--hop", type=float, default=0.5, help="每秒几帧 chroma")
a = ap.parse_args()

NFFT = 1024; HOP = int(a.sr / a.hop)

def pcm(path, ss=None, t=None):
    cmd = ["ffmpeg", "-v", "error"]
    if ss: cmd += ["-ss", str(ss)]
    if t: cmd += ["-t", str(t)]
    cmd += ["-i", path, "-ac", "1", "-ar", str(a.sr), "-f", "s16le", "-"]
    raw = subprocess.run(cmd, capture_output=True).stdout
    n = len(raw) // 2
    return struct.unpack("<%dh" % n, raw[:n*2])

def fft(x):
    n = len(x)
    if n == 1: return x
    e = fft(x[0::2]); o = fft(x[1::2])
    out = [0j]*n
    for k in range(n//2):
        t = cmath.exp(-2j*math.pi*k/n) * o[k]
        out[k] = e[k]+t; out[k+n//2] = e[k]-t
    return out

# 预计算频点→音级
WIN = [0.5-0.5*math.cos(2*math.pi*i/(NFFT-1)) for i in range(NFFT)]
PC = []
for k in range(NFFT//2+1):
    f = k*a.sr/NFFT
    PC.append(-1 if (f < 65 or f > 4000) else round(69+12*math.log2(f/440.0)) % 12)

def chroma(x):
    T = max(0, (len(x)-NFFT)//HOP)
    out = []
    for ti in range(T):
        off = ti*HOP
        seg = x[off:off+NFFT]
        sp = fft([complex(seg[i]*WIN[i]) for i in range(NFFT)])
        C = [0.0]*12
        for k in range(NFFT//2+1):
            p = PC[k]
            if p >= 0: C[p] += abs(sp[k])
        C = [math.log1p(v) for v in C]
        m = sorted(C)[6]
        out.append([v-m for v in C])
    return out

def norm(v):
    n = math.sqrt(sum(x*x for x in v)) + 1e-9
    return [x/n for x in v]

def best_match(q, ref):
    """q: 查询 chroma 列表; ref: 参考序列; 返回 (最佳相关系数, 位置帧号)"""
    if not q or not ref or len(ref) <= len(q): return None, None
    qn = [norm(c) for c in q]
    rn = [norm(c) for c in ref]
    best = (-2, -1)
    L = len(qn)
    for s in range(0, len(rn)-L):
        acc = 0.0
        for i in range(L):
            a1, b1 = qn[i], rn[s+i]
            acc += sum(a1[k]*b1[k] for k in range(12))
        if acc > best[0]: best = (acc, s)
    return best

print(f"=== 视频取样: {a.vstart}s 起 {a.vlen}s ===")
qv = chroma(pcm(a.video, a.vstart, a.vlen))
print(f"  查询 chroma 帧数 = {len(qv)}")

rows = []
for p in sorted(glob.glob(os.path.join(a.seg_dir, "avseq*.mpg"))):
    ref = chroma(pcm(p))
    corr, pos = best_match(qv, ref)
    if corr is None:
        rows.append((os.path.basename(p), None, None)); continue
    rows.append((os.path.basename(p), corr, pos/a.hop if pos is not None else None))
    print(f"  {os.path.basename(p):<14} 最佳相关={corr:+.3f}  位置≈{pos/a.hop:.1f}s")

rows.sort(key=lambda r: (r[1] is None, -(r[1] or -9)))
print("\n=== 排序 ===")
for name, c, pos in rows:
    print(f"  {name:<14} corr={'%.3f'%c if c is not None else 'n/a':>7}  位置={('%.1fs'%pos) if pos else '-'}")
if rows and rows[0][1] is not None:
    print(f"\n→ 最可能：{rows[0][0]}（在 ≈{rows[0][2]:.1f}s 处）")
