Archive DLL-free personalized HRTF binaural renderer
This commit is contained in:
@@ -0,0 +1,243 @@
|
||||
import importlib.util
|
||||
import sys
|
||||
import tempfile
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
|
||||
ROOT = Path(__file__).resolve().parents[1]
|
||||
SRC = ROOT / "src"
|
||||
if str(SRC) not in sys.path:
|
||||
sys.path.insert(0, str(SRC))
|
||||
if str(ROOT) not in sys.path:
|
||||
sys.path.insert(0, str(ROOT))
|
||||
|
||||
import main
|
||||
from adm_atmos import q_to_adm_xyz
|
||||
from binaural_metadata import OamdPositionTimeline
|
||||
from binaural_renderer import (
|
||||
DEFAULT_PERSONALIZED_HEADPHONE,
|
||||
RosellaBinauralRenderer,
|
||||
resolve_personalized_headphone,
|
||||
)
|
||||
from oamd_bits import q_of
|
||||
from rosella_direct import BINAURAL_PROFILE_NAMES, direct_and_room_send, special_lfe_direct
|
||||
from rosella_filterbank import HybridAnalysis, HybridSynthesis, QmfAnalysis, QmfSynthesis
|
||||
from rosella_model import load_personalized_headphone
|
||||
from speaker_wav import BinauralPcmSpool, write_pcm_wav
|
||||
|
||||
_MODEL_CANDIDATES = list(
|
||||
(ROOT / "tests" / "binauraltests" / "evidence").glob(
|
||||
"*/test.personalized_headphone"))
|
||||
TEST_MODEL = _MODEL_CANDIDATES[0] if _MODEL_CANDIDATES else Path()
|
||||
|
||||
|
||||
class BinauralProductionTest(unittest.TestCase):
|
||||
def test_cli_exposes_only_near_mid_far_and_defaults_mid_float32(self):
|
||||
parser = main.build_parser()
|
||||
args = parser.parse_args(["input.eac3", "--binaural"])
|
||||
self.assertEqual(args.binaural_mode, "mid")
|
||||
self.assertEqual(args.binaural_format, "float32")
|
||||
action = next(a for a in parser._actions if a.dest == "binaural_mode")
|
||||
self.assertEqual(tuple(action.choices), ("near", "mid", "far"))
|
||||
self.assertNotIn("off", BINAURAL_PROFILE_NAMES)
|
||||
|
||||
def test_default_model_path_and_missing_model_message(self):
|
||||
self.assertEqual(
|
||||
DEFAULT_PERSONALIZED_HEADPHONE,
|
||||
ROOT / "HRTF" / "binaural.personalized_headphone")
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
missing = Path(directory) / "missing.personalized_headphone"
|
||||
with self.assertRaisesRegex(FileNotFoundError, "未找到双耳模型"):
|
||||
resolve_personalized_headphone(missing)
|
||||
|
||||
def test_runtime_filterbanks_are_double_precision(self):
|
||||
qmf = QmfAnalysis(1)
|
||||
hybrid = HybridAnalysis(1)
|
||||
hybrid_synthesis = HybridSynthesis(1)
|
||||
qmf_synthesis = QmfSynthesis(1)
|
||||
self.assertEqual(qmf.coefficients.dtype, np.float64)
|
||||
self.assertEqual(qmf.history.dtype, np.float64)
|
||||
self.assertEqual(hybrid.low_kernel.dtype, np.float64)
|
||||
self.assertEqual(hybrid.high_history.dtype, np.complex128)
|
||||
self.assertEqual(qmf_synthesis.basis.dtype, np.float64)
|
||||
source = np.zeros((1, 1, 64), dtype=np.float64)
|
||||
q = qmf.process_chunk(source)
|
||||
h = hybrid.process_chunk(q)
|
||||
self.assertEqual(q.dtype, np.complex128)
|
||||
self.assertEqual(h.dtype, np.complex128)
|
||||
back = hybrid_synthesis.process_chunk(h)
|
||||
time = qmf_synthesis.process_chunk(back)
|
||||
self.assertEqual(back.dtype, np.complex128)
|
||||
self.assertEqual(time.dtype, np.float64)
|
||||
|
||||
@unittest.skipUnless(
|
||||
(ROOT / "tests" / "binauraltests" / "hqmf.py").is_file(),
|
||||
"research filterbank reference not present")
|
||||
def test_compact_kernel_archive_matches_research_float64_filterbanks(self):
|
||||
reference_root = ROOT / "tests" / "binauraltests"
|
||||
spec = importlib.util.spec_from_file_location(
|
||||
"binaural_reference_hqmf", reference_root / "hqmf.py")
|
||||
reference = importlib.util.module_from_spec(spec)
|
||||
spec.loader.exec_module(reference)
|
||||
evidence = reference_root / "evidence"
|
||||
source = np.fromfile(
|
||||
evidence / "hqmf_kernel" / "hqmf_validation_input.f32le",
|
||||
dtype="<f4").astype(np.float64).reshape(-1, 1, 64)
|
||||
production_qmf = QmfAnalysis(1).process_chunk(source)
|
||||
reference_qmf = reference.QmfAnalysisFast(
|
||||
evidence / "hqmf_kernel" / "hqmf_analysis_manifest.json",
|
||||
1, dtype=np.float64).process_chunk(source)
|
||||
np.testing.assert_array_equal(production_qmf, reference_qmf)
|
||||
|
||||
production_hybrid = HybridAnalysis(1).process_chunk(production_qmf)
|
||||
reference_hybrid = reference.HybridAnalysis(
|
||||
evidence / "hybrid_kernel" / "hybrid_analysis_manifest.json",
|
||||
1, dtype=np.float64).process_chunk(reference_qmf)
|
||||
np.testing.assert_array_equal(production_hybrid, reference_hybrid)
|
||||
|
||||
production_qmf_back = HybridSynthesis(1).process_chunk(production_hybrid)
|
||||
reference_qmf_back = reference.HybridSynthesis(
|
||||
evidence / "hybrid_synthesis_kernel" / "hybrid_synthesis_manifest.json",
|
||||
1, dtype=np.float64).process_chunk(reference_hybrid)
|
||||
np.testing.assert_array_equal(production_qmf_back, reference_qmf_back)
|
||||
|
||||
production_time = QmfSynthesis(1).process_chunk(production_qmf_back)
|
||||
reference_time = reference.QmfSynthesisFast(
|
||||
evidence / "qmf_synthesis_kernel" / "qmf_synthesis_manifest.json",
|
||||
1, dtype=np.float64).process_chunk(reference_qmf_back)
|
||||
np.testing.assert_array_equal(production_time, reference_time)
|
||||
|
||||
def test_special_lfe_is_fixed_16_band_complex128_without_room_send(self):
|
||||
result = special_lfe_direct()
|
||||
self.assertEqual(result.gains.dtype, np.complex128)
|
||||
np.testing.assert_array_equal(result.gains[0], result.gains[1])
|
||||
self.assertTrue(np.any(result.gains[:, :16] != 0))
|
||||
np.testing.assert_array_equal(result.gains[:, 16:], 0)
|
||||
self.assertEqual(float(result.room_send), 0.0)
|
||||
|
||||
def test_oamd_timing_retains_outer_block_delay_and_ramp(self):
|
||||
timeline = OamdPositionTimeline()
|
||||
initial = {
|
||||
"values": {
|
||||
(1, "q1"): q_of(0, 62),
|
||||
(1, "q2"): q_of(0, 62),
|
||||
(1, "q3"): q_of(0, 15),
|
||||
},
|
||||
"block_offset_samples": 0,
|
||||
"ramp_duration_samples": 1536,
|
||||
}
|
||||
timeline.submit_update(initial, frame_start_sample=0)
|
||||
old = np.asarray(q_to_adm_xyz(q_of(0, 62), q_of(0, 62), q_of(0, 15)))
|
||||
np.testing.assert_allclose(timeline.positions_at(0)[0], old)
|
||||
|
||||
target_q1 = q_of(62, 62)
|
||||
changed = {
|
||||
"values": {(1, "q1"): target_q1},
|
||||
"block_offset_samples": 32,
|
||||
"ramp_duration_samples": 1536,
|
||||
}
|
||||
timeline.submit_update(
|
||||
changed,
|
||||
frame_start_sample=1536,
|
||||
outer_sample_offset=16,
|
||||
object_delay_samples=1473,
|
||||
)
|
||||
start = 1536 + 16 + 32 + 1473 + 64
|
||||
duration = 1536 - 64
|
||||
target = np.asarray(q_to_adm_xyz(target_q1, q_of(0, 62), q_of(0, 15)))
|
||||
np.testing.assert_allclose(timeline.positions_at(start - 1)[0], old)
|
||||
np.testing.assert_allclose(timeline.positions_at(start)[0], old)
|
||||
np.testing.assert_allclose(
|
||||
timeline.positions_at(start + duration // 2)[0],
|
||||
old + (target - old) * 0.5,
|
||||
)
|
||||
np.testing.assert_allclose(timeline.positions_at(start + duration)[0], target)
|
||||
|
||||
@unittest.skipUnless(TEST_MODEL.is_file(), "test model not present")
|
||||
def test_model_parser_and_direct_path_promote_to_double(self):
|
||||
model = load_personalized_headphone(TEST_MODEL)
|
||||
self.assertEqual(model.sample_rate, 48000)
|
||||
self.assertEqual(len(model.coefficients), 16033)
|
||||
self.assertEqual(
|
||||
model.coefficient_sha256,
|
||||
"2d4b40c27925ec8827556585d90c371384458d31bb84b20f94a946145558ea27",
|
||||
)
|
||||
result = direct_and_room_send(model, (0.0, 1.0, 0.0), 3)
|
||||
self.assertEqual(result.gains.dtype, np.complex128)
|
||||
with self.assertRaisesRegex(ValueError, "near, mid, or far"):
|
||||
direct_and_room_send(model, (0.0, 1.0, 0.0), 0)
|
||||
|
||||
@unittest.skipUnless(TEST_MODEL.is_file(), "test model not present")
|
||||
def test_renderer_keeps_float64_state_and_compensates_961_samples(self):
|
||||
renderer = RosellaBinauralRenderer(
|
||||
TEST_MODEL,
|
||||
chunk_frames=1,
|
||||
tail_seconds=0,
|
||||
room_impulse_slots=32,
|
||||
)
|
||||
source = np.zeros((1536, 16), dtype=np.float32)
|
||||
source[0, 1] = 1.0
|
||||
first = renderer.render_frame(source)
|
||||
tail = renderer.finish()
|
||||
self.assertEqual(first.dtype, np.float64)
|
||||
self.assertEqual(tail.dtype, np.float64)
|
||||
self.assertEqual(len(first), 1536 - 961)
|
||||
self.assertEqual(renderer.qmf_analysis.history.dtype, np.float64)
|
||||
self.assertEqual(renderer.hybrid_analysis.high_history.dtype, np.complex128)
|
||||
self.assertEqual(renderer.core.gains.dtype, np.complex128)
|
||||
self.assertEqual(renderer.core.room.tail.dtype, np.complex128)
|
||||
|
||||
@unittest.skipUnless(TEST_MODEL.is_file(), "test model not present")
|
||||
def test_native_backend_matches_python_backend(self):
|
||||
source = np.zeros((1536, 16), dtype=np.float32)
|
||||
source[0, 0] = 0.1
|
||||
source[64, 1] = -0.2
|
||||
python_renderer = RosellaBinauralRenderer(
|
||||
TEST_MODEL, backend="python", chunk_frames=1,
|
||||
tail_seconds=0, room_impulse_slots=64)
|
||||
try:
|
||||
native_renderer = RosellaBinauralRenderer(
|
||||
TEST_MODEL, backend="native",
|
||||
native_library=ROOT / "lib" / "eac3joc_core.dll",
|
||||
chunk_frames=1, tail_seconds=0)
|
||||
except (OSError, RuntimeError):
|
||||
python_renderer.close()
|
||||
self.skipTest("native binaural backend not built")
|
||||
try:
|
||||
self.assertEqual(native_renderer.dsp_backend, "native")
|
||||
python_output = np.concatenate((
|
||||
python_renderer.render_frame(source), python_renderer.finish()))
|
||||
native_output = np.concatenate((
|
||||
native_renderer.render_frame(source), native_renderer.finish()))
|
||||
np.testing.assert_allclose(
|
||||
native_output, python_output, rtol=0.0, atol=2.0e-14)
|
||||
finally:
|
||||
python_renderer.close()
|
||||
native_renderer.close()
|
||||
|
||||
def test_binaural_spool_preserves_float64_then_trims_tail(self):
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
raw = Path(directory) / "binaural.raw"
|
||||
wav = Path(directory) / "binaural.wav"
|
||||
spool = BinauralPcmSpool(raw, 8, tail_threshold=1.0e-8)
|
||||
values = np.asarray([
|
||||
[0.0, 0.0],
|
||||
[0.25, -0.25],
|
||||
[1.0 + 2.0 ** -40, 0.0],
|
||||
[1.0e-9, 0.0],
|
||||
], dtype=np.float64)
|
||||
spool.write_frame(values)
|
||||
spool.finalize(minimum_samples=2)
|
||||
self.assertEqual(spool.values.dtype, np.float64)
|
||||
self.assertEqual(spool.sample_count, 3)
|
||||
self.assertEqual(spool.clipped_values, 1)
|
||||
info = write_pcm_wav(wav, spool.values, "float32")
|
||||
self.assertEqual(info["sample_count"], 3)
|
||||
self.assertEqual(info["channel_count"], 2)
|
||||
spool.close()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user