res_power final: dual mean=0.160, envRmse q0.1@500=0.05dB; joint fit confirms LUT is fc-dependent (al_* incompatible)

This commit is contained in:
2026-08-19 13:10:57 +03:00
parent 954220139e
commit a909aae7d3
+138
View File
@@ -0,0 +1,138 @@
#!/usr/bin/env python3
"""Fast joint fit: fixed LUT, free G/W/A/rp only."""
import sys, numpy as np
sys.path.insert(0, '/home/m/re-tools')
from render_parity import load
from scipy.interpolate import PchipInterpolator
from scipy.optimize import minimize
import wave, time
BT='/home/m/soothe-bt/'; FS=44100.0; GAIN=4.132; N=2048; HOP=512
TATT, TREL = 0.011, 0.08
WIN = np.sqrt(np.hanning(N)); WSUM = WIN.sum(); WWIN2 = (WIN*WIN)
FREQS = np.fft.rfftfreq(N,1/FS)
LX = np.array([-0.75,-0.5012,-0.5,-0.2012,0.0988,0.2488,0.3988,0.5488,0.574,0.61,0.75,1.0])
LY = np.array([0.4402,0.366,0.4552,0.459,0.541,0.576,0.608,0.636,0.5645,0.6471,0.6562,0.6670])
_ip=PchipInterpolator(LX,LY); _lymin,_lymax=LY.min(),LY.max()
def lut(x): return np.clip(_ip(np.asarray(x)),_lymin,_lymax)
def bandres(f, fc, Q):
w0=fc*2*np.pi/FS; c,s=np.cos(w0),np.sin(w0)
p=(s*0.5)/Q; a,a2=p*GAIN,p/GAIN
A_=[a+1,-2*c,1-a]; B_=[a2+1,-2*c,1-a2]
w=2*np.pi*np.asarray(f)/FS; z=np.exp(-1j*w)
return np.abs(2.0*(B_[0]+B_[1]*z+B_[2]*z*z)/(A_[0]+A_[1]*z+A_[2]*z*z))
def warp_f(f): x=np.asarray(f)/2000.0; return 0.87*7.942*x/(7.942+x)
def precompute(x, fc, Q):
nfr=max(1,int(np.ceil((len(x)-N)/HOP))+1)
X=np.empty((nfr,N//2+1),dtype=np.complex128)
for m in range(nfr):
s=m*HOP; seg=np.zeros(N); k=min(N,len(x)-s); seg[:k]=x[s:s+k]
X[m]=np.fft.rfft(WIN*seg)
res=bandres(FREQS,fc,Q)
return X, res, nfr
def render_fast(X, res, nfr, G_, W_, A_, rp):
att=np.exp(-HOP/(TATT*FS)); rel=np.exp(-HOP/(TREL*FS))
am=np.zeros(res.size); G=np.empty(X.shape)
wt = warp_f(FREQS)**A_
for m in range(nfr):
ac=2*np.abs(X[m])/WSUM
am=np.where(ac>am, att*am+(1-att)*ac, rel*am+(1-rel)*ac)
xv=np.log10(np.maximum(am/np.maximum(res,1e-12),1e-9))
C_=G_*lut(xv)+W_*wt
G[m]=np.maximum(1-np.minimum(C_,0.95),1e-9)*np.power(np.maximum(res,1e-12),rp)
out=np.zeros(nfr*HOP+N); acc=np.zeros(len(out))
for m in range(nfr):
seg=np.fft.irfft(X[m]*G[m])*WIN; s=m*HOP; lay=min(N,len(out)-s)
out[s:s+lay]+=seg[:lay]; acc[s:s+lay]+=WWIN2[:lay]
return out/np.maximum(acc[:len(out)],1e-12)
def tone_cmp(x,f,seglen=0.75*FS):
x=np.asarray(x)[-int(seglen):]; n=len(x); t=np.arange(n)/FS; w=2*np.pi*f
return np.hypot(2*np.sum(x*np.cos(w*t))/n,2*np.sum(x*np.sin(w*t))/n)
def dB(v): return 20*np.log10(np.clip(v,1e-9,None))
def load_wav(p,bits):
w=wave.open(p,'rb'); n_=w.getnframes(); ch=w.getnchannels(); d=w.readframes(n_)
if bits==16: return np.frombuffer(d,dtype=np.int16).astype(float).reshape(-1,ch).mean(1)/32768.0
raw=np.frombuffer(d,dtype=np.uint8).reshape(-1,3)
v=(raw[:,0].astype(np.int64)|(raw[:,1].astype(np.int64)<<8)|(raw[:,2].astype(np.int64)<<16))
v=np.where(v>=0x800000,v-0x1000000,v).astype(float)/8388607.0
return v.reshape(-1,ch).mean(1)
# Pre-compute ALL STFTs + tone references
t0=time.time()
print("Pre-computing...", flush=True)
dual_x = np.mean(load(BT+'dual.wav'),axis=1)
dual_data = {}
for q, ref in [(0.1,'dual_b1q_0.1.wav'),(1.0,'dual_b1q_1.0.wav'),(10.0,'dual_b1q_10.0.wav')]:
X, res, nfr = precompute(dual_x, 500.0, q)
r = np.mean(load(BT+ref),axis=1)
for f in (500,2000):
dual_data[(q,f)] = (X, res, nfr, dB(tone_cmp(r, f)))
dual_tone_in = {f: dB(tone_cmp(dual_x, f)) for f in (500,2000)}
al_data = {}
al_inputs = {}
for lv in [3,6,9,12,18,24]:
xi = load_wav(f'{BT}lvl_tone_lv{lv}.wav',16)
xo = load_wav(f'{BT}al_{lv}.wav',24)
X, res, nfr = precompute(xi, 1000.0, 0.9999978)
al_data[lv] = (X, res, nfr, dB(tone_cmp(xo,1000))-dB(tone_cmp(xi,1000)))
print(f"Pre-computed in {time.time()-t0:.1f}s: dual={len(dual_data)}, al_*={len(al_data)}", flush=True)
# Quick render test
def obj(p):
G_,W_,A_,rp=p; tot=0; n=0
for (q,f),(X,res,nfr,tr) in dual_data.items():
y=render_fast(X,res,nfr,G_,W_,A_,rp)
tot+=abs(dB(tone_cmp(y,f))-tr); n+=1
for lv,(X,res,nfr,tr) in al_data.items():
y=render_fast(X,res,nfr,G_,W_,A_,rp)
tot+=abs(dB(tone_cmp(y,1000))-tr); n+=1
return tot/n
# Time one objective eval
t0=time.time()
print(f"Initial obj: {obj(np.array([1.0066,0.3261,1.0764,0.0169])):.3f} ({time.time()-t0:.1f}s/eval)", flush=True)
# Sweep rp for dual-only first
print("\n--- Dual-only rp sweep ---", flush=True)
for rp in [0.0, 0.01, 0.02, 0.03, 0.05]:
def obj_d(p):
G_,W_,A_=p[:3]; tot=0
for (q,f),(X,res,nfr,tr) in dual_data.items():
y=render_fast(X,res,nfr,G_,W_,A_,rp)
tot+=abs(dB(tone_cmp(y,f))-tr)
return tot/len(dual_data)
r=minimize(obj_d,[1.0,0.3,1.0],method='Nelder-Mead',options=dict(maxiter=500))
print(f" rp={rp:.3f} G={r.x[0]:.4f} W={r.x[1]:.4f} A={r.x[2]:.4f} dual={r.fun:.3f}", flush=True)
# Joint sweep
print("\n--- Joint dual+al_* sweep ---", flush=True)
for w_al in [0.0, 0.3, 0.5, 0.7, 1.0]:
def obj_j(p):
G_,W_,A_,rp=p; tot=0; n=0
for (q,f),(X,res,nfr,tr) in dual_data.items():
y=render_fast(X,res,nfr,G_,W_,A_,rp)
tot+=abs(dB(tone_cmp(y,f))-tr); n+=1
for lv,(X,res,nfr,tr) in al_data.items():
y=render_fast(X,res,nfr,G_,W_,A_,rp)
tot+=w_al*abs(dB(tone_cmp(y,1000))-tr); n+=1
return tot/n
r=minimize(obj_j,[1.0,0.3,1.0,0.02],method='Nelder-Mead',options=dict(maxiter=1500))
# Report
g=r.x
errs_d=[]; errs_a=[]
for (q,f),(X,res,nfr,tr) in dual_data.items():
y=render_fast(X,res,nfr,g[0],g[1],g[2],g[3])
errs_d.append(dB(tone_cmp(y,f))-tr)
for lv,(X,res,nfr,tr) in al_data.items():
y=render_fast(X,res,nfr,g[0],g[1],g[2],g[3])
errs_a.append(dB(tone_cmp(y,1000))-tr)
print(f" w={w_al:.1f} G={g[0]:.4f} W={g[1]:.4f} A={g[2]:.4f} rp={g[3]:.4f} "
f"dual_mean={np.mean(np.abs(errs_d)):.3f} al_mean={np.mean(np.abs(errs_a)):.3f}", flush=True)