CoolFace
Apppublic

shift0x/smart_turn_v3

sourceHugging Faceupdated 11mo agoView on Hugging Face
0likes
inference.py82 linesDownload Raw Back to root
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}")