#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""YIN 基频估计（去八度误差）：累积均值归一化差值函数 + 绝对门限 + 抛物线插值。
相比自相关取全局峰，YIN 取"第一个低于门限的谷"，可避免强二次谐波导致的降八度错误。
用法: yin-f0.py <vocals.wav> <标签> [--min 80] [--max 700]
"""
from __future__ import annotations
import argparse
import numpy as np
import soundfile as sf

FRAME, HOP = 1024, 256      # 16kHz: 64ms / 16ms
THRESH = 0.15               # YIN 绝对门限
RMS_FLOOR = 0.03            # 相对全局 RMS 的能量门限（滤掉纯伴奏/静音段）


def yin_frame(fr: np.ndarray, sr: int, tau_min: int, tau_max: int):
    W = len(fr)
    tau_max = min(tau_max, W - 1)
    # 差值函数
    d = np.empty(tau_max + 1)
    d[0] = 0.0
    for tau in range(1, tau_max + 1):
        diff = fr[:W - tau] - fr[tau:]
        d[tau] = float(np.dot(diff, diff))
    # 累积均值归一化
    cmnd = np.ones(tau_max + 1)
    run = 0.0
    for tau in range(1, tau_max + 1):
        run += d[tau]
        cmnd[tau] = d[tau] * tau / run if run > 0 else 1.0
    # 第一个低于门限的局部谷
    tau = -1
    for t in range(tau_min, tau_max):
        if cmnd[t] < THRESH:
            while t + 1 <= tau_max and cmnd[t + 1] < cmnd[t]:
                t += 1
            tau = t
            break
    if tau < 0:
        t = int(np.argmin(cmnd[tau_min:tau_max + 1])) + tau_min
        if cmnd[t] > 0.35:
            return None
        tau = t
    # 抛物线插值
    if tau_min < tau < tau_max:
        a, b, c = cmnd[tau - 1], cmnd[tau], cmnd[tau + 1]
        denom = 2 * (2 * b - a - c)
        if denom != 0:
            tau = tau + (c - a) / denom
    return sr / tau


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("wav"); ap.add_argument("tag")
    ap.add_argument("--min", type=float, default=80.0)
    ap.add_argument("--max", type=float, default=700.0)
    a = ap.parse_args()

    x, sr = sf.read(a.wav, always_2d=True)
    mono = x.mean(axis=1).astype(np.float64)
    if sr != 16000:
        idx = np.linspace(0, len(mono) - 1, int(len(mono) * 16000 / sr))
        mono = np.interp(idx, np.arange(len(mono)), mono)
        sr = 16000
    tau_min, tau_max = int(sr / a.max), int(sr / a.min)
    win = np.hanning(FRAME)
    grms = float(np.sqrt((mono ** 2).mean())) + 1e-12
    f0s = []
    for s in range(0, len(mono) - FRAME, HOP):
        fr = mono[s:s + FRAME]
        if float(np.sqrt((fr ** 2).mean())) < RMS_FLOOR * grms:
            continue
        fr = (fr - fr.mean()) * win
        f = yin_frame(fr, sr, tau_min, tau_max)
        if f:
            f0s.append(f)
    f0 = np.array(f0s)
    if f0.size < 20:
        print(f"{a.tag}: 有效清音帧过少({f0.size})"); return
    q = np.percentile(f0, [5, 25, 50, 75, 95])
    names = ['C', 'C#', 'D', 'D#', 'E', 'F', 'F#', 'G', 'G#', 'A', 'A#', 'B']
    def nn(f):
        m = int(round(69 + 12 * np.log2(f / 440.0)))
        return f"{names[m % 12]}{m // 12 - 1}"
    print(f"{a.tag}: 有效清音帧 {f0.size} (总帧数约 {len(mono)//HOP})")
    print(f"  F0 分位 5/25/50/75/95 = {q[0]:.0f} / {q[1]:.0f} / {q[2]:.0f} / {q[3]:.0f} / {q[4]:.0f} Hz")
    print(f"  中位 F0 = {q[2]:.0f} Hz ≈ {nn(q[2])}   音域(5-95%) {nn(q[0])} ~ {nn(q[4])}")
    band = "女声区(≥175Hz)" if q[2] >= 175 else ("中间区(155-175)" if q[2] > 155 else "男声区(≤155Hz)")
    print(f"  判定: {band}   [参考: 男声典型 100-155 / 女声典型 175-260 / 童声 260+]")


if __name__ == "__main__":
    main()
