session parameter required in OnnxWrapper
This commit is contained in:
parent
2431a6bf91
commit
62444ce746
1 changed files with 5 additions and 27 deletions
|
|
@ -54,7 +54,7 @@ class OnnxWrapper():
|
||||||
ONNX Runtime wrapper for Silero VAD model with per-instance state.
|
ONNX Runtime wrapper for Silero VAD model with per-instance state.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, session: OnnxSession = None, force_onnx_cpu=False):
|
def __init__(self, session: OnnxSession, force_onnx_cpu=False):
|
||||||
self._shared_session = session
|
self._shared_session = session
|
||||||
self.sample_rates = session.sample_rates
|
self.sample_rates = session.sample_rates
|
||||||
self.reset_states()
|
self.reset_states()
|
||||||
|
|
@ -313,10 +313,8 @@ class FixedVADIterator(VADIterator):
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
# Test JIT model
|
# vad = FixedVADIterator(load_jit_vad())
|
||||||
print("Testing JIT model...")
|
vad = FixedVADIterator(OnnxWrapper(session=load_onnx_session()))
|
||||||
model = load_jit_vad()
|
|
||||||
vad = FixedVADIterator(model)
|
|
||||||
|
|
||||||
audio_buffer = np.array([0] * 512, dtype=np.float32)
|
audio_buffer = np.array([0] * 512, dtype=np.float32)
|
||||||
result = vad(audio_buffer)
|
result = vad(audio_buffer)
|
||||||
|
|
@ -326,23 +324,3 @@ if __name__ == "__main__":
|
||||||
audio_buffer = np.array([0] * 511, dtype=np.float32)
|
audio_buffer = np.array([0] * 511, dtype=np.float32)
|
||||||
result = vad(audio_buffer)
|
result = vad(audio_buffer)
|
||||||
print(f" 511 samples: {result}")
|
print(f" 511 samples: {result}")
|
||||||
|
|
||||||
# Test ONNX with shared session
|
|
||||||
print("\nTesting ONNX with shared session...")
|
|
||||||
shared_session = load_onnx_session()
|
|
||||||
|
|
||||||
# Create two independent VAD iterators sharing the same session
|
|
||||||
vad1 = FixedVADIterator(OnnxWrapper(session=shared_session))
|
|
||||||
vad2 = FixedVADIterator(OnnxWrapper(session=shared_session))
|
|
||||||
|
|
||||||
# Both should work independently
|
|
||||||
audio_buffer = np.array([0] * 512, dtype=np.float32)
|
|
||||||
result1 = vad1(audio_buffer)
|
|
||||||
result2 = vad2(audio_buffer)
|
|
||||||
print(f" VAD1 result: {result1}")
|
|
||||||
print(f" VAD2 result: {result2}")
|
|
||||||
|
|
||||||
# Verify they have separate states
|
|
||||||
print(f" VAD1 and VAD2 share session: {vad1.model._shared_session is vad2.model._shared_session}")
|
|
||||||
print(f" VAD1 and VAD2 have separate state: {vad1.model._state is not vad2.model._state}")
|
|
||||||
print("\nAll tests passed!")
|
|
||||||
Loading…
Reference in a new issue