From 954220139ef2e8906a255e1966aa85db8cf16072 Mon Sep 17 00:00:00 2001 From: Matiq Date: Wed, 19 Aug 2026 11:32:25 +0300 Subject: [PATCH] =?UTF-8?q?res=5Fpower=20dual-only:=20mean=3D0.080,=20envR?= =?UTF-8?q?mse=200.64=E2=86=920.05=20at=20q0.1@500;=20G/W/A/rp=3D1.0066/0.?= =?UTF-8?q?3261/1.0764/0.0169?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- handoff/joint_fast.py | 128 ++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 128 insertions(+) create mode 100644 handoff/joint_fast.py diff --git a/handoff/joint_fast.py b/handoff/joint_fast.py new file mode 100644 index 0000000..14d19c2 --- /dev/null +++ b/handoff/joint_fast.py @@ -0,0 +1,128 @@ +#!/usr/bin/env python3 +"""Fast joint refit: pre-compute STFTs, only recompute gain+OLA per iteration.""" +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 + +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() + +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): x=np.asarray(f)/2000.0; return 0.87*7.942*x/(7.942+x) + +def precompute(x, fc, Q): + """Pre-compute STFT + res + freqs (independent of G/W/A/rp).""" + 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) + freqs=np.fft.rfftfreq(N,1/FS); res=bandres(freqs,fc,Q) + return X, freqs, res, nfr + +def render_fast(X, freqs, res, nfr, G_, W_, A_, rp): + att=np.exp(-HOP/(TATT*FS)); rel=np.exp(-HOP/(TREL*FS)) + am=np.zeros(freqs.size); G=np.empty(X.shape) + freq_term = warp(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_*freq_term + 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]+=(WIN*WIN)[: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 +print("Pre-computing STFTs...") +dual_x = np.mean(load(BT+'dual.wav'),axis=1) +dual_data = {} # (q, f) -> (X, res, nfr, tone_ref) +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, freqs, res, nfr = precompute(dual_x, 500.0, q) + r = np.mean(load(BT+ref),axis=1) + for f in (500,2000): + tone_ref = dB(tone_cmp(r, f)) + dual_data[(q,f)] = (X, res, nfr, tone_ref) +dual_tone_in = {f: dB(tone_cmp(dual_x, f)) for f in (500,2000)} + +al_data = {} # lv -> (X, res, nfr, ref_ratio) +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, freqs, res, nfr = precompute(xi, 1000.0, 0.9999978) + ref_ratio = tone_cmp(xo,1000)/tone_cmp(xi,1000) + al_data[lv] = (X, res, nfr, ref_ratio) +print(f"Dual: {len(dual_data)} cases, al_*: {len(al_data)} cases") + +# Joint objective +def obj(logp, w_al=0.5): + lG,W_,A_,rp = logp; G_=np.exp(lG) + errs=[] + # Dual (weight 1.0) + for (q,f),(X,res,nfr,tone_ref) in dual_data.items(): + y=render_fast(X,freqs,res,nfr,G_,W_,A_,rp) + out=dB(tone_cmp(y,f)) + errs.append(out-tone_ref) + # al_* (weight w_al) + for lv,(X,res,nfr,ref_ratio) in al_data.items(): + y=render_fast(X,freqs,res,nfr,G_,W_,A_,rp) + out_ratio=tone_cmp(y,1000) + out_dB=dB(out_ratio) + ref_dB=dB(ref_ratio) + errs.append(w_al*(out_dB-ref_dB)) + return np.mean(np.abs(errs)) + +best=(999,None,None) +for w_al in [0.0, 0.5, 1.0]: + for p0 in [np.log([0.97,0.35,1.09,0.05]), np.log([1.0,0.3,1.5,0.1])]: + r = minimize(obj, p0, args=(w_al,), method='Nelder-Mead', options=dict(maxiter=3000, xatol=1e-5)) + if r.fun < best[0]: best=(r.fun, r.x, w_al) + +p=best[1]; G_=np.exp(p[0]); w_al=best[2] +print(f'\nBEST: G={G_:.4f} W={p[1]:.4f} A={p[2]:.4f} rp={p[3]:.4f} w_al={w_al:.1f} mean={best[0]:.3f}') + +print('\ndual:') +for q in [0.1,1.0,10.0]: + for f in [500,2000]: + X,res,nfr,tone_ref=dual_data[(q,f)] + y=render_fast(X,freqs,res,nfr,G_,p[1],p[2],p[3]) + err=dB(tone_cmp(y,f))-tone_ref + print(f' q{q} @{f}: err={err:+.2f} ref={tone_ref-dual_tone_in[f]:+.2f} out={dB(tone_cmp(y,f))-dual_tone_in[f]:+.2f}') + +print('\nal_*:') +for lv in [3,6,9,12,18,24]: + X,res,nfr,ref_ratio=al_data[lv] + y=render_fast(X,freqs,res,nfr,G_,p[1],p[2],p[3]) + out_dB=dB(tone_cmp(y,1000)); ref_dB=dB(ref_ratio) + print(f' lv{lv}: ref={ref_dB:+.2f} out={out_dB:+.2f} err={out_dB-ref_dB:+.2f}') \ No newline at end of file