shift0x/smart_turn_v3
0
1import numpy as np2import onnxruntime as ort3from transformers import WhisperFeatureExtractor4 5ONNX_MODEL_PATH = "smart-turn-v3.0.onnx"6 7def build_session(onnx_path):8 so = ort.SessionOptions()9 so.execution_mode = ort.ExecutionMode.ORT_SEQUENTIAL10 so.inter_op_num_threads = 111 so.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL12 return ort.InferenceSession(onnx_path, sess_options=so)13 14feature_extractor = WhisperFeatureExtractor(chunk_length=8)15session = build_session(ONNX_MODEL_PATH)16 17def truncate_audio_to_last_n_seconds(audio_array, n_seconds=8, sample_rate=16000):18 """Truncate audio to last n seconds or pad with zeros to meet n seconds."""19 max_samples = n_seconds * sample_rate20 if len(audio_array) > max_samples:21 return audio_array[-max_samples:]22 elif len(audio_array) < max_samples:23 # Pad with zeros at the beginning24 padding = max_samples - len(audio_array)25 return np.pad(audio_array, (padding, 0), mode='constant', constant_values=0)26 return audio_array27 28 29def predict_endpoint(audio_array):30 """31 Predict whether an audio segment is complete (turn ended) or incomplete.32 33 Args:34 audio_array: Numpy array containing audio samples at 16kHz35 36 Returns:37 Dictionary containing prediction results:38 - prediction: 1 for complete, 0 for incomplete39 - probability: Probability of completion (sigmoid output)40 """41 42 # Truncate to 8 seconds (keeping the end) or pad to 8 seconds43 audio_array = truncate_audio_to_last_n_seconds(audio_array, n_seconds=8)44 45 # Process audio using Whisper's feature extractor46 inputs = feature_extractor(47 audio_array,48 sampling_rate=16000,49 return_tensors="np",50 padding="max_length",51 max_length=8 * 16000,52 truncation=True,53 do_normalize=True,54 )55 56 # Extract features and ensure correct shape for ONNX57 input_features = inputs.input_features.squeeze(0).astype(np.float32)58 input_features = np.expand_dims(input_features, axis=0) # Add batch dimension59 60 # Run ONNX inference61 outputs = session.run(None, {"input_features": input_features})62 63 # Extract probability (ONNX model returns sigmoid probabilities)64 probability = outputs[0][0].item()65 66 # Make prediction (1 for Complete, 0 for Incomplete)67 prediction = 1 if probability > 0.5 else 068 69 return {70 "prediction": prediction,71 "probability": probability,72 }73 74 75# Example usage76if __name__ == "__main__":77 # Create a dummy audio array for testing (1 second of random audio)78 dummy_audio = np.random.randn(16000).astype(np.float32)79 80 result = predict_endpoint(dummy_audio)81 print(f"Prediction: {result['prediction']}")82 print(f"Probability: {result['probability']:.4f}")