Download ir_source_localizer.py from zeechimp/ir-source-localizer: direct link, hf CLI and curl.
- Browser
- Download file 23.7 kB
-
https://huggingface.co/zeechimp/ir-source-localizer/resolve/main/ir_source_localizer.py
- Command line
-
hf download hf://zeechimp/ir-source-localizer/ir_source_localizer.py
-
curl -L -o ir_source_localizer.py https://huggingface.co/zeechimp/ir-source-localizer/resolve/main/ir_source_localizer.py
23.7 kB
| #!/usr/bin/env python3 | |
| """ | |
| ir_source_localizer.py | |
| ====================== | |
| Localize an acoustic source from its impulse response at known | |
| microphones. RANSAC over receiver subsets rejects outlying | |
| arrival-time detections. | |
| Fix history (5 iterations to get here): | |
| v1 wide-window argmax -- locks on strong reflection | |
| v2 first local max after rise -- picks leading-edge noise | |
| v3 half-peak leading edge -- smoothing merges direct+reflection | |
| v4 squared signal + tight smoothing -- r2 still +0.7 ms biased | |
| v5 quality-weighted loss -- weighting helped nothing | |
| v6 RANSAC over receiver subsets -- rejects the bad receiver | |
| The insight: three of four receivers detect the direct within | |
| 0.02 ms. One is off by 0.70 ms. Any scheme that averages all four | |
| inherits the outlier. RANSAC finds the consensus of three and | |
| drops the fourth. | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import itertools | |
| import json | |
| import math | |
| import wave | |
| from dataclasses import dataclass, asdict | |
| from typing import List, Optional, Sequence, Tuple | |
| import numpy as np | |
| SPEED_OF_SOUND = 343.0 | |
| # ===================================================================== | |
| # §1 Result type | |
| # ===================================================================== | |
| class LocalizationResult: | |
| x: float | |
| y: float | |
| sigma_x: float | |
| sigma_y: float | |
| detected_times: List[float] | |
| detection_quality: List[float] | |
| predicted_times: List[float] | |
| arrival_residual_ms: List[float] | |
| rms_residual_ms: float | |
| n_receivers: int | |
| n_inliers: int | |
| inlier_mask: List[bool] | |
| method: str = "energy envelope + RANSAC + Adam" | |
| def to_dict(self) -> dict: | |
| return asdict(self) | |
| def __str__(self) -> str: | |
| return (f"source at ({self.x:+.3f}, {self.y:+.3f}) m " | |
| f"± ({self.sigma_x*100:.1f}, " | |
| f"{self.sigma_y*100:.1f}) cm " | |
| f"inliers {self.n_inliers}/{self.n_receivers} " | |
| f"rms resid {self.rms_residual_ms:.3f} ms") | |
| # ===================================================================== | |
| # §2 Direct arrival detection (unchanged from v4) | |
| # ===================================================================== | |
| def detect_direct_arrival(ir, sample_rate, sigma_guess, | |
| threshold=0.30): | |
| dt = 1.0 / sample_rate | |
| w = max(3, int(round(1.0 * sigma_guess / dt))) | |
| kernel = np.ones(w) / w | |
| env = np.convolve(ir ** 2, kernel, mode="same") | |
| peak = float(env.max()) | |
| if peak < 1e-12: | |
| return None, 0.0 | |
| above = env >= threshold * peak | |
| if not above.any(): | |
| return None, 0.0 | |
| first_idx = int(np.argmax(above)) | |
| lo = first_idx | |
| hi = min(len(env), first_idx + 3 * w) | |
| if hi - lo < 3: | |
| return None, 0.0 | |
| peak_idx = lo + int(np.argmax(env[lo:hi])) | |
| return float(peak_idx * dt), float(env[peak_idx] / peak) | |
| def detect_all(irs, sample_rate, sigma_guess): | |
| times, qualities = [], [] | |
| for ir in irs: | |
| t, q = detect_direct_arrival(ir, sample_rate, sigma_guess) | |
| times.append(t) | |
| qualities.append(q) | |
| return times, qualities | |
| # ===================================================================== | |
| # §3 Position estimation | |
| # ===================================================================== | |
| def predict_arrivals(sx, sy, sz, receivers, c=SPEED_OF_SOUND): | |
| d = np.sqrt((receivers[:, 0] - sx) ** 2 | |
| + (receivers[:, 1] - sy) ** 2 | |
| + (receivers[:, 2] - sz) ** 2) | |
| return d / c | |
| def arrival_loss(sx, sy, sz, receivers, times): | |
| pred = predict_arrivals(sx, sy, sz, receivers) | |
| return float(np.mean((pred - times) ** 2)) | |
| def localize_grid(times, receivers, sz, x_range, y_range, n_grid=60): | |
| xs = np.linspace(x_range[0], x_range[1], n_grid) | |
| ys = np.linspace(y_range[0], y_range[1], n_grid) | |
| best = (float(xs[0]), float(ys[0]), float("inf")) | |
| for i in range(n_grid): | |
| for j in range(n_grid): | |
| sx = float(xs[i]); sy = float(ys[j]) | |
| d = np.sqrt((receivers[:, 0] - sx) ** 2 | |
| + (receivers[:, 1] - sy) ** 2 | |
| + (receivers[:, 2] - sz) ** 2) | |
| t_pred = d / SPEED_OF_SOUND | |
| err = float(np.mean((t_pred - times) ** 2)) | |
| if err < best[2]: | |
| best = (sx, sy, err) | |
| return best | |
| def refine_position(sx0, sy0, sz, receivers, times, | |
| n_iter=300, lr=0.02): | |
| p = np.array([sx0, sy0]) | |
| def L(p): | |
| return arrival_loss(p[0], p[1], sz, receivers, times) | |
| m = np.zeros(2); v = np.zeros(2) | |
| best_p = p.copy(); best_L = L(p) | |
| eps = 1e-3 | |
| for t in range(1, n_iter + 1): | |
| g = np.zeros(2) | |
| for i in range(2): | |
| rp = p.copy(); rp[i] += eps | |
| rm = p.copy(); rm[i] -= eps | |
| g[i] = (L(rp) - L(rm)) / (2.0 * eps) | |
| m = 0.9 * m + 0.1 * g | |
| v = 0.999 * v + 0.001 * g * g | |
| mh = m / (1 - 0.9 ** t) | |
| vh = v / (1 - 0.999 ** t) | |
| p = p - lr * mh / (np.sqrt(vh) + 1e-8) | |
| Li = L(p) | |
| if Li < best_L: | |
| best_L = Li | |
| best_p = p.copy() | |
| return float(best_p[0]), float(best_p[1]), best_L | |
| # ===================================================================== | |
| # §4 RANSAC over receiver subsets -- the v6 fix | |
| # ===================================================================== | |
| def ransac_localize(times, receivers, sz, | |
| x_range, y_range, | |
| inlier_threshold_ms=0.30, | |
| n_grid=50, | |
| n_refine_ransac=100, | |
| n_refine_final=300, | |
| seed=0): | |
| """Try every subset of 3 receivers. Count how many of the K | |
| receivers agree with each solution within `inlier_threshold_ms`. | |
| Return the position with the largest consensus, then refine on | |
| the inlier set only. | |
| If K == 3, there is one subset and the result is the ordinary | |
| solution. For K >= 4, this rejects outlier receivers. | |
| """ | |
| K = len(times) | |
| if K < 3: | |
| raise ValueError("need at least 3 receivers") | |
| threshold_s = inlier_threshold_ms / 1000.0 | |
| subsets = list(itertools.combinations(range(K), 3)) | |
| best = None # (n_inliers, inlier_indices, sx, sy) | |
| for sub in subsets: | |
| sub_idx = np.array(sub) | |
| recv_sub = receivers[sub_idx] | |
| times_sub = np.asarray(times)[sub_idx] | |
| # Coarse grid on the subset | |
| sx0, sy0, _ = localize_grid(times_sub, recv_sub, sz, | |
| x_range, y_range, n_grid) | |
| # Fine on the subset | |
| sx, sy, _ = refine_position(sx0, sy0, sz, recv_sub, | |
| times_sub, | |
| n_iter=n_refine_ransac, lr=0.03) | |
| # Count inliers among all receivers | |
| pred = predict_arrivals(sx, sy, sz, receivers) | |
| resid = np.abs(pred - np.asarray(times)) | |
| inliers = np.where(resid <= threshold_s)[0] | |
| score = len(inliers) | |
| if best is None or score > best[0]: | |
| best = (score, inliers, sx, sy) | |
| _, inlier_indices, sx, sy = best | |
| # Final refinement on inliers only | |
| inliers_arr = np.asarray(inlier_indices) | |
| if len(inliers_arr) >= 3: | |
| sx, sy, _ = refine_position( | |
| sx, sy, sz, | |
| receivers[inliers_arr], | |
| np.asarray(times)[inliers_arr], | |
| n_iter=n_refine_final, lr=0.02) | |
| return sx, sy, inlier_indices | |
| # ===================================================================== | |
| # §5 Uncertainty (Monte Carlo over inliers) | |
| # ===================================================================== | |
| def estimate_uncertainty(sx, sy, sz, receivers, times, inliers, | |
| x_range, y_range, | |
| sigma_t, n_samples=50, n_grid=40, seed=0): | |
| rng = np.random.default_rng(seed) | |
| R = receivers[inliers] | |
| T = np.asarray(times)[inliers] | |
| xs = np.empty(n_samples); ys = np.empty(n_samples) | |
| for k in range(n_samples): | |
| tp = T + rng.normal(0, sigma_t, size=len(T)) | |
| sx0, sy0, _ = localize_grid(tp, R, sz, x_range, y_range, n_grid) | |
| xr, yr, _ = refine_position(sx0, sy0, sz, R, tp, | |
| n_iter=120, lr=0.03) | |
| xs[k] = xr; ys[k] = yr | |
| return float(np.std(xs)), float(np.std(ys)) | |
| # ===================================================================== | |
| # §6 Main API | |
| # ===================================================================== | |
| def localize(irs, sample_rate, receivers, sz, | |
| x_range, y_range, | |
| sigma_guess=4e-4, | |
| n_grid=60, n_refine=300, | |
| estimate_unc=True, n_unc_samples=50, | |
| unc_sigma_t=None, | |
| inlier_threshold_ms=0.30, | |
| seed=0): | |
| if len(irs) != len(receivers): | |
| raise ValueError( | |
| f"got {len(irs)} IRs and {len(receivers)} receivers") | |
| if len(irs) < 3: | |
| raise ValueError(f"need at least 3 receivers, got {len(irs)}") | |
| receivers_arr = np.asarray(receivers, dtype=np.float64) | |
| times, qualities = detect_all(irs, sample_rate, sigma_guess) | |
| bad = [k for k, t in enumerate(times) if t is None] | |
| if bad: | |
| raise RuntimeError( | |
| f"detection failed at receivers {bad}") | |
| times_arr = np.asarray(times, dtype=np.float64) | |
| qualities = np.asarray(qualities, dtype=np.float64) | |
| sx, sy, inliers = ransac_localize( | |
| times_arr, receivers_arr, sz, | |
| x_range, y_range, | |
| inlier_threshold_ms=inlier_threshold_ms, | |
| n_grid=50, | |
| n_refine_ransac=100, | |
| n_refine_final=n_refine, | |
| seed=seed) | |
| t_pred = predict_arrivals(sx, sy, sz, receivers_arr) | |
| resid_ms = (t_pred - times_arr) * 1000.0 | |
| rms_ms = float(np.sqrt(np.mean(resid_ms ** 2))) | |
| sigma_x = float("nan"); sigma_y = float("nan") | |
| if estimate_unc: | |
| sigma_t = (unc_sigma_t if unc_sigma_t is not None | |
| else 0.3 * sigma_guess) | |
| sigma_x, sigma_y = estimate_uncertainty( | |
| sx, sy, sz, receivers_arr, times_arr, inliers, | |
| x_range, y_range, sigma_t, | |
| n_samples=n_unc_samples, seed=seed) | |
| inlier_mask = [bool(k in set(int(i) for i in inliers)) | |
| for k in range(len(irs))] | |
| return LocalizationResult( | |
| x=sx, y=sy, | |
| sigma_x=sigma_x, sigma_y=sigma_y, | |
| detected_times=[float(t) for t in times], | |
| detection_quality=[float(q) for q in qualities], | |
| predicted_times=[float(t) for t in t_pred], | |
| arrival_residual_ms=[float(r) for r in resid_ms], | |
| rms_residual_ms=rms_ms, | |
| n_receivers=len(irs), | |
| n_inliers=len(inliers), | |
| inlier_mask=inlier_mask, | |
| ) | |
| # ===================================================================== | |
| # §7 WAV I/O | |
| # ===================================================================== | |
| def load_wav(path): | |
| with wave.open(path, "rb") as w: | |
| n_ch = w.getnchannels() | |
| samp_w = w.getsampwidth() | |
| sr = w.getframerate() | |
| raw = w.readframes(w.getnframes()) | |
| if samp_w == 2: | |
| data = (np.frombuffer(raw, dtype=np.int16) | |
| .astype(np.float64) / 32768.0) | |
| elif samp_w == 4: | |
| data = (np.frombuffer(raw, dtype=np.int32) | |
| .astype(np.float64) / 2147483648.0) | |
| elif samp_w == 1: | |
| data = ((np.frombuffer(raw, dtype=np.uint8) | |
| .astype(np.float64) - 128.0) / 128.0) | |
| else: | |
| raise ValueError(f"unsupported sample width: {samp_w}") | |
| if n_ch > 1: | |
| data = data.reshape(-1, n_ch).mean(axis=1) | |
| return data, sr | |
| # ===================================================================== | |
| # §8 Simulator | |
| # ===================================================================== | |
| def _image_sources(sx, sy, sz, Lx, Ly, Lz, n_reflect): | |
| pos, n_r = [], [] | |
| for i in range(-n_reflect, n_reflect + 1): | |
| for j in range(-n_reflect, n_reflect + 1): | |
| for k in range(-n_reflect, n_reflect + 1): | |
| nr = abs(i) + abs(j) + abs(k) | |
| if nr > n_reflect: | |
| continue | |
| for xs_ in (sx + 2 * i * Lx, -sx + 2 * i * Lx): | |
| for ys_ in (sy + 2 * j * Ly, -sy + 2 * j * Ly): | |
| for zs_ in (sz + 2 * k * Lz, | |
| -sz + 2 * k * Lz): | |
| pos.append((xs_, ys_, zs_)) | |
| n_r.append(nr) | |
| return (np.array(pos, dtype=np.float64), | |
| np.array(n_r, dtype=np.float64)) | |
| def _pulse_mex(delta, sigma): | |
| x = delta / sigma | |
| return (1.0 - x * x) * np.exp(-x * x / 2.0) | |
| def simulate_ir(receiver, source, room, times, | |
| alpha=0.72, sigma=4e-4, n_reflect=2): | |
| Lx, Ly, Lz = room | |
| sx, sy, sz = source | |
| pos, nr = _image_sources(sx, sy, sz, Lx, Ly, Lz, n_reflect) | |
| d = np.linalg.norm(pos - receiver[None, :], axis=1) | |
| t_arr = d / SPEED_OF_SOUND | |
| amp = (alpha ** nr) / np.maximum(d, 0.15) | |
| delta = times[:, None] - t_arr[None, :] | |
| return (amp[None, :] * _pulse_mex(delta, sigma)).sum(axis=1) | |
| def simulate_all(receivers, source, room, times, | |
| alpha=0.72, sigma=4e-4, n_reflect=2, | |
| noise_frac=0.0, seed=0): | |
| irs = np.stack([ | |
| simulate_ir(r, source, room, times, alpha, sigma, n_reflect) | |
| for r in receivers | |
| ]) | |
| if noise_frac > 0: | |
| rng = np.random.default_rng(seed) | |
| rms = float(np.sqrt(np.mean(irs ** 2))) | |
| irs = irs + noise_frac * rms * rng.standard_normal(irs.shape) | |
| return irs | |
| # ===================================================================== | |
| # §9 Self-test | |
| # ===================================================================== | |
| def self_test(verbose=True): | |
| checks = [] | |
| sr = 48000 | |
| T = 4000 | |
| times = np.arange(T) / sr | |
| sigma = 4e-4 | |
| pulse_time = 0.010 | |
| pulse = _pulse_mex(times - pulse_time, sigma) | |
| t_det, q = detect_direct_arrival(pulse, sr, sigma) | |
| checks.append(("clean pulse within 0.5 ms", | |
| t_det is not None | |
| and abs(t_det - pulse_time) < 5e-4)) | |
| checks.append(("quality near 1.0", q > 0.9)) | |
| rng = np.random.default_rng(0) | |
| noisy = pulse + 0.05 * rng.standard_normal(T) | |
| t_det2, _ = detect_direct_arrival(noisy, sr, sigma) | |
| checks.append(("5% noise within 1 ms", | |
| t_det2 is not None | |
| and abs(t_det2 - pulse_time) < 1e-3)) | |
| room = (5.0, 4.0, 2.8) | |
| source = (1.5, 1.5, 1.0) | |
| recv = np.array([ | |
| [4.5, 3.5, 1.0], [0.5, 3.5, 1.0], | |
| [4.5, 0.5, 1.0], [0.5, 0.5, 1.0], | |
| ]) | |
| # RANSAC self-test: inject an outlier time. | |
| t_arr = predict_arrivals(source[0], source[1], source[2], recv) | |
| t_with_outlier = t_arr.copy() | |
| t_with_outlier[2] += 0.7e-3 # add 0.7 ms to receiver 2 | |
| sx, sy, inliers = ransac_localize( | |
| t_with_outlier, recv, source[2], | |
| (0.2, 4.8), (0.2, 3.8), | |
| inlier_threshold_ms=0.3, | |
| n_grid=50, n_refine_ransac=100, n_refine_final=200, | |
| seed=0) | |
| err = math.hypot(sx - source[0], sy - source[1]) | |
| checks.append((f"RANSAC rejects outlier (err {err*100:.2f} cm)", | |
| err < 0.05)) | |
| checks.append(("outlier receiver is not an inlier", | |
| 2 not in inliers)) | |
| # End-to-end at 2% noise | |
| T_ir = 2000 | |
| times_ir = np.arange(T_ir) / sr | |
| irs = simulate_all(recv, source, room, times_ir, | |
| alpha=0.72, sigma=sigma, n_reflect=2, | |
| noise_frac=0.02, seed=1) | |
| res = localize(irs, sr, recv, sz=source[2], | |
| x_range=(0.2, 4.8), y_range=(0.2, 3.8), | |
| sigma_guess=sigma, n_grid=50, n_refine=250, | |
| estimate_unc=False) | |
| err_e2e = math.hypot(res.x - source[0], res.y - source[1]) | |
| checks.append((f"end-to-end within 5 cm (err {err_e2e*100:.2f})", | |
| err_e2e < 0.05)) | |
| # Uncertainty | |
| res_u = localize(irs, sr, recv, sz=source[2], | |
| x_range=(0.2, 4.8), y_range=(0.2, 3.8), | |
| sigma_guess=sigma, n_grid=40, n_refine=200, | |
| estimate_unc=True, n_unc_samples=20, seed=3) | |
| checks.append(("uncertainty positive", | |
| res_u.sigma_x > 0 and res_u.sigma_y > 0)) | |
| try: | |
| localize(irs[:2], sr, recv[:2], sz=source[2], | |
| x_range=(0, 5), y_range=(0, 4)) | |
| checks.append(("rejects <3 receivers", False)) | |
| except ValueError: | |
| checks.append(("rejects <3 receivers", True)) | |
| try: | |
| import tempfile, os | |
| with tempfile.NamedTemporaryFile(suffix=".wav", | |
| delete=False) as f: | |
| path = f.name | |
| sig = np.clip(irs[0], -1.0, 1.0) | |
| with wave.open(path, "wb") as wf: | |
| wf.setnchannels(1); wf.setsampwidth(2) | |
| wf.setframerate(sr) | |
| wf.writeframes((sig * 32767).astype(np.int16).tobytes()) | |
| loaded, sr_loaded = load_wav(path) | |
| os.unlink(path) | |
| checks.append(("WAV sr preserved", sr_loaded == sr)) | |
| checks.append(("WAV len preserved", len(loaded) == len(irs[0]))) | |
| checks.append(("WAV values close", | |
| float(np.max(np.abs(loaded - sig))) < 1e-3)) | |
| except Exception as e: | |
| checks.append((f"WAV failed: {e}", False)) | |
| passed = sum(1 for _, ok in checks if ok) | |
| if verbose: | |
| print() | |
| print("=" * 74) | |
| print("SELF-TEST") | |
| print("=" * 74) | |
| for name, ok in checks: | |
| mark = "PASS" if ok else "FAIL" | |
| print(f" [{mark}] {name}") | |
| print() | |
| print(f" {passed}/{len(checks)} correct") | |
| return passed, len(checks) | |
| # ===================================================================== | |
| # §10 Demo | |
| # ===================================================================== | |
| def banner(t, w=76): | |
| print() | |
| print("=" * w) | |
| print(t) | |
| print("=" * w) | |
| def demo(): | |
| sr = 48000 | |
| sigma = 4e-4 | |
| room = (5.0, 4.0, 2.8) | |
| source_true = (1.5, 1.5, 1.0) | |
| recv = np.array([ | |
| [4.5, 3.5, 1.0], [0.5, 3.5, 1.0], | |
| [4.5, 0.5, 1.0], [0.5, 0.5, 1.0], | |
| ]) | |
| banner("IR SOURCE LOCALIZER v6 -- RANSAC outlier rejection") | |
| print(f" true source : " | |
| f"({source_true[0]:.2f}, {source_true[1]:.2f}, " | |
| f"{source_true[2]:.2f}) m") | |
| print(f" room : {room[0]} x {room[1]} x {room[2]} m") | |
| print() | |
| print(f" {'noise':>6} {'x':>8} {'y':>8} {'err cm':>8} " | |
| f"{'inliers':>8} {'rms ms':>8}") | |
| print(" " + "-" * 60) | |
| times_ir = np.arange(2000) / sr | |
| for noise in (0.005, 0.01, 0.02, 0.05, 0.10): | |
| irs = simulate_all(recv, source_true, room, times_ir, | |
| noise_frac=noise, seed=1) | |
| try: | |
| res = localize(irs, sr, recv, sz=source_true[2], | |
| x_range=(0.2, 4.8), y_range=(0.2, 3.8), | |
| sigma_guess=sigma, n_grid=50, | |
| n_refine=250, estimate_unc=False) | |
| err_cm = math.hypot(res.x - source_true[0], | |
| res.y - source_true[1]) * 100 | |
| print(f" {noise*100:>5.1f}% " | |
| f"{res.x:>8.3f} {res.y:>8.3f} {err_cm:>8.2f} " | |
| f"{res.n_inliers:>3}/{res.n_receivers:<4} " | |
| f"{res.rms_residual_ms:>8.3f}") | |
| except RuntimeError as e: | |
| print(f" {noise*100:>5.1f}% fail: {e}") | |
| banner("DETAIL -- 2% noise, single run") | |
| irs = simulate_all(recv, source_true, room, times_ir, | |
| noise_frac=0.02, seed=1) | |
| res = localize(irs, sr, recv, sz=source_true[2], | |
| x_range=(0.2, 4.8), y_range=(0.2, 3.8), | |
| sigma_guess=sigma, n_grid=50, n_refine=250, | |
| estimate_unc=True, n_unc_samples=40, seed=7) | |
| print(f" {res}") | |
| print() | |
| print(f" {'recv':>5} {'detected (ms)':>14} " | |
| f"{'predicted (ms)':>15} {'resid (ms)':>11} " | |
| f"{'quality':>8} {'inlier':>7}") | |
| print(" " + "-" * 72) | |
| for k in range(res.n_receivers): | |
| print(f" {k:>5} {res.detected_times[k]*1000:>14.4f} " | |
| f"{res.predicted_times[k]*1000:>15.4f} " | |
| f"{res.arrival_residual_ms[k]:>+11.4f} " | |
| f"{res.detection_quality[k]:>8.3f} " | |
| f"{str(res.inlier_mask[k]):>7}") | |
| print() | |
| print(f" RMS arrival residual : {res.rms_residual_ms:.4f} ms") | |
| print(f" position uncertainty : " | |
| f"sx = {res.sigma_x*100:.2f} cm, " | |
| f"sy = {res.sigma_y*100:.2f} cm") | |
| # ===================================================================== | |
| # §11 CLI | |
| # ===================================================================== | |
| def _parse_xyz(s): | |
| parts = s.replace(",", " ").split() | |
| if len(parts) != 3: | |
| raise argparse.ArgumentTypeError("expected 3 numbers") | |
| return tuple(float(p) for p in parts) | |
| def main(): | |
| p = argparse.ArgumentParser() | |
| sub = p.add_subparsers(dest="cmd") | |
| sub.add_parser("self-test") | |
| sub.add_parser("demo") | |
| p_loc = sub.add_parser("localize") | |
| p_loc.add_argument("wavs", nargs="+") | |
| p_loc.add_argument("--sample-rate", type=int, required=True) | |
| p_loc.add_argument("--receivers", nargs="+", required=True, | |
| type=_parse_xyz) | |
| p_loc.add_argument("--sz", type=float, required=True) | |
| p_loc.add_argument("--x-range", nargs=2, type=float, required=True) | |
| p_loc.add_argument("--y-range", nargs=2, type=float, required=True) | |
| p_loc.add_argument("--sigma", type=float, default=4e-4) | |
| p_loc.add_argument("--grid", type=int, default=50) | |
| p_loc.add_argument("--refine", type=int, default=300) | |
| p_loc.add_argument("--inlier-threshold-ms", type=float, default=0.30) | |
| p_loc.add_argument("--no-unc", action="store_true") | |
| p_loc.add_argument("--json", action="store_true") | |
| args = p.parse_args() | |
| if args.cmd in (None, "self-test"): | |
| self_test() | |
| if args.cmd is None: | |
| print() | |
| print(" usage: python ir_source_localizer.py [self-test|demo|localize ...]") | |
| return | |
| if args.cmd == "demo": | |
| demo() | |
| return | |
| if args.cmd == "localize": | |
| irs = [] | |
| for path in args.wavs: | |
| data, sr = load_wav(path) | |
| if sr != args.sample_rate: | |
| raise SystemExit(f"sample rate mismatch: {path}") | |
| irs.append(data) | |
| res = localize(irs, args.sample_rate, args.receivers, | |
| sz=args.sz, | |
| x_range=tuple(args.x_range), | |
| y_range=tuple(args.y_range), | |
| sigma_guess=args.sigma, | |
| n_grid=args.grid, n_refine=args.refine, | |
| estimate_unc=not args.no_unc, | |
| inlier_threshold_ms=args.inlier_threshold_ms) | |
| if args.json: | |
| print(json.dumps(res.to_dict(), indent=2)) | |
| else: | |
| print(res) | |
| if __name__ == "__main__": | |
| main() |