#!/usr/bin/env python3 """ fit_vlaw_params.py — Fit VLAW parameters (α, β, c, Δ, γ₀) per configuration group. VLAW model (framed_model.cpp:198-200): cs = α * log1p(lvl / β) + c + (delta ? Δ : 0) applied_gain = 10^(-γ₀ * cs / 20) Currently hardcoded for dual(q=0.5): α=3.2193, β=0.4927, c=0.5423, Δ=7.46-0.5423, γ₀=1.79 Need to fit these for each (fc, q, sens) configuration group: t1kq: fc=800..1200, q=1.0, sens=12 t1k: fc=500..2000, q=1.0, sens=12 al: fc=1000, q=1.0, sens=3..24 res: fc=300..700, q=1.0, sens=12 dual: fc=500, q=0.1..10.0, sens=12 """ import numpy as np import json import os import sys import subprocess sys.path.insert(0, '/home/m/re-tools/scripts') import corpus corpus.RB = '/home/m/re-tools/dsp/build/render48k' def structural_cases(): out = [] for name, inp, args, ref, f in corpus.build_cases(): joined = [','.join(args)] if len(args) == 3 else args out.append((name, inp, joined, ref, f)) return out def group_key(name): return name.split('_')[0] def load_ref_errors(): """Load baseline_bridge.json for target errors.""" with open('scripts/baseline_bridge.json') as f: return json.load(f) def render_vlaw(inp, out, args, alpha, beta, c, delta, gamma0, extra_env=None): """Run render48k with VLAW parameters.""" env = { **os.environ, 'RT_VLAW': '1', 'RT_VLAW_ALPHA': str(alpha), 'RT_VLAW_BETA': str(beta), 'RT_VLAW_C': str(c), 'RT_VLAW_DELTA': str(delta), 'RT_VLAW_GAMMA0': str(gamma0), 'RT_SYN': '1', 'RT_NOWARP': '1', 'RT_NOIIR3': '1', 'RT_IIR12': '0', } if extra_env: env.update(extra_env) subprocess.run( [corpus.RB, inp, out] + args, capture_output=True, text=True, env=env, cwd='/home/m/re-tools' ) def eval_config(alpha, beta, c, delta, gamma0, cases_subset=None): """Evaluate VLAW params on cases, return per-group mean abs error.""" all_cases = structural_cases() if cases_subset: all_cases = [c for c in all_cases if group_key(c[0]) in cases_subset] refs = load_ref_errors() errs = {} for name, inp, args, ref, f in all_cases: out = f'/tmp/vlaw_fit_{name}.wav' render_vlaw(inp, out, args, alpha, beta, c, delta, gamma0) if not os.path.exists(out) or os.path.getsize(out) == 0: errs[name] = 999.0 continue try: ref_sig = corpus.load_mono(ref) out_sig = corpus.load_mono(out) min_len = min(len(ref_sig), len(out_sig)) ref_sig = ref_sig[-min_len:] out_sig = out_sig[-min_len:] ref_ta = corpus.ta(ref_sig, f) out_ta = corpus.ta(out_sig, f) err_db = corpus.db(out_ta / ref_ta) errs[name] = err_db except Exception as e: print(f"Error on {name}: {e}") errs[name] = 999.0 # Group stats groups = {} for k, v in errs.items(): g = group_key(k) groups.setdefault(g, []).append(v) out = {g: float(np.mean(np.abs(v))) for g, v in groups.items()} out['TOTAL'] = float(np.mean(np.abs(list(errs.values())))) return out, errs def fit_alpha_beta_c(cases_to_fit): """Coordinate descent on (α, β, c) for a specific case group.""" # For now, grid search best = None best_err = float('inf') # Search ranges around current dual(q=0.5) values for alpha in np.linspace(2.5, 4.0, 8): for beta in np.linspace(0.3, 0.7, 8): for c in np.linspace(0.2, 1.0, 8): stats, _ = eval_config(alpha, beta, c, 6.9, 1.79, cases_to_fit) total = stats['TOTAL'] if total < best_err: best_err = total best = (alpha, beta, c, stats) print(f" New best: α={alpha:.4f}, β={beta:.4f}, c={c:.4f}, TOTAL={total:.4f}") return best def main(): # Build case map by group all_cases = structural_cases() groups = {} for name, inp, args, ref, f in all_cases: g = group_key(name) groups.setdefault(g, []).append(name) print("Available groups:", list(groups.keys())) for g, names in groups.items(): print(f" {g}: {len(names)} cases") # Start with dual group (already calibrated) print("\n=== Testing dual(q=0.5) baseline ===") stats, errs = eval_config(3.2193, 0.4927, 0.5423, 6.9177, 1.79, ['dual']) print(f"Dual stats: {stats}") # Now fit for each group for g in ['t1kq', 't1k', 'al', 'res', 'dual']: if g not in groups: continue print(f"\n=== Fitting {g} ===") best = fit_alpha_beta_c([g]) if best: alpha, beta, c, stats = best print(f" {g} best: α={alpha:.4f}, β={beta:.4f}, c={c:.4f}") print(f" Stats: {stats}") if __name__ == '__main__': main()