244 lines
11 KiB
Python
244 lines
11 KiB
Python
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()
|