silero-vad

要检测媒体中的语音,可以使用的方法比较多,这里我们采用深度学习的方法,使用一个开源的silero-vad

根据源代码主页上的说法,silero-vad在6000多种语言上进行过训练,准确率很高,另外它的模型文件却很小,JIT模型文件仅仅有2兆左右大小。

silero-vad的使用也非常简单,使用Python的silero-vad库的示例如下:

from silero_vad import load_silero_vad, read_audio, get_speech_timestamps
model = load_silero_vad()
wav = read_audio('path_to_audio_file')
speech_timestamps = get_speech_timestamps(
  wav,
  model,
  return_seconds=True,  # Return speech timestamps in seconds (default is samples)
)

或者使用torch.hub:

import torch
torch.set_num_threads(1)

model, utils = torch.hub.load(repo_or_dir='snakers4/silero-vad', model='silero_vad')
(get_speech_timestamps, _, read_audio, _, _) = utils

wav = read_audio('path_to_audio_file')
speech_timestamps = get_speech_timestamps(
  wav,
  model,
  return_seconds=True,  # Return speech timestamps in seconds (default is samples)
)

需要注意的是:silero_vad使用了libsox.so这个动态链接库解析音频,所以测试以上代码需要安装sox

onnxruntime

但是,在flutter应用中,就没有现成的silero-vad模块可以使用了。

好在模型是现成的。要在flutter应用中使用深度学习模型进行推理,我们可以使用onnxruntime,直接调用模型文件进行推理 。

onnx是一种开放的模型格式,很多深度学习库都支持。silero-vad开源的模型文件也是onnx格式的。我们使用flutter中的onnxruntime-flutter插件,再加上silero-vad的onnx格式模型文件,就可以了。

加入onnxruntime

我们使用flutter add onnxruntime把onnxruntime加入当前项目。

之后将在pubspec.yaml中看到类似:

onnxruntime: ^1.4.1

加入模型文件

hugging face 官网搜索silero-vad,找到社区版的模型文件的目录,下载第一个model.onnx。

之后,把model.onnx放入项目的asserts/models目录。

OrtSession加载模型

创建一个类VoiceDetectionService,使用OrtSession的fromBuffer方法,加载模型。

import 'package:onnxruntime/onnxruntime.dart';

  Future<void> loadModel() async {
    try {
      final modelData = await rootBundle.load('assets/models/model.onnx');
      final options = OrtSessionOptions();
      OrtSession? _session = OrtSession.fromBuffer(modelData.buffer.asUint8List(), options);
    } catch (e) {
      debugPrint('Failed to load VAD model: $e');
    }
  }

OrtSession执行推理

加载模型以后,就可以使用OrtSession的run方法,执行推理。

根据silero_vad.utils_vad.py的源码:

       if sr in [8000, 16000]:  
           ort_inputs = {'input': x.numpy(), 'state': self._state.numpy(), 'sr': np.array(sr, dtype='int64  
')}  
           ort_outs = self.session.run(None, ort_inputs)  
           out, state = ort_outs  
           self._state = torch.from_numpy(state)  
       else:  
           raise ValueError()  
  
       self._context = x[..., -context_size:]  
       self._last_sr = sr  
       self._last_batch_size = batch_size  
  
       out = torch.from_numpy(out)  
       return out

silero-vad的这个模型的输入参数是一个哈西:分别是input,state与sr。

其中state是采样点的序列值,如果声音频率为16000,则为512个。如果为8000,则为256个。

state是每次返回的状态,sr则是一个表示采样频率的值。

以上代码使用dart实现,则为:

      final inputTensor = OrtValueTensor.createTensorWithDataList(
          modelInput, [1, windowSize]);
      final srTensor = OrtValueTensor.createTensorWithData(options.sampleRate);

      final inputs = {
        'input': inputTensor as OrtValue,
        'state': state,
        'sr': srTensor
      };

      final runOptions = OrtRunOptions();
      final outputs = _session!.run(runOptions, inputs);
      state = outputs[1]! as OrtValueTensor;

      inputTensor.release();
      srTensor.release();

StreamTransformer

因为从ffmpeg解析出来的是流式数据,而我们执行推理的是固定数量的采样点,所以可以使用Dart中的StreamTransformer,做一个转换。

import 'dart:async';  
import 'dart:typed_data';  
  
class WindowingTransformer  
    implements StreamTransformer<List<int>, Float32List> {  
  final int windowSize;  
  final int bytesPerSample;  
  
  WindowingTransformer({  
    required this.windowSize,  
    this.bytesPerSample = 4,  
  });  
  
    
  Stream<Float32List> bind(Stream<List<int>> stream) {  
    final controller = StreamController<Float32List>();  
    List<int> buffer = [];  
    final windowSizeBytes = windowSize * bytesPerSample;  
  
    controller.onListen = () {  
      final subscription = stream.listen(  
        (chunk) {  
          buffer.addAll(chunk);  
          while (buffer.length >= windowSizeBytes) {  
            final windowBytes = buffer.sublist(0, windowSizeBytes);  
            controller  
                .add(Uint8List.fromList(windowBytes).buffer.asFloat32List());  
            buffer = buffer.sublist(windowSizeBytes);  
          }  
        },  
        onError: controller.addError,  
        onDone: () {  
          controller.close();  
        },  
      );  
  
      controller.onCancel = () {  
        return subscription.cancel();  
      };  
    };  
  
    return controller.stream;  
  }  
  
    
  StreamTransformer<RS, RT> cast<RS, RT>() => StreamTransformer.castFrom(this);  
}

总结

最终,我们使用前文介绍的ffmpeg,把音频文件解析到一个管道,之后通过上文的Transformer,最后调用OrtSession的run执行推理,判断出每段音频是否是人声,整个工作就结束了。


import 'dart:async';

import 'package:fcsplayer/app/services/ffmpeg_processing_service.dart';
import 'package:fcsplayer/app/utils/stream_transformers.dart';
import 'package:flutter/foundation.dart';
import 'package:flutter/services.dart';
import 'package:onnxruntime/onnxruntime.dart';

class VoiceDetectionService {
  OrtSession? _session;
  final int sampleRate = 16000;
  final int windowSize = 512;
  final FFmpegProcessingService _ffmpegProcessingService =
      FFmpegProcessingService();

  VoiceDetectionService();

  Future<void> loadModel() async {
    if (_session != null) return;
    try {
      final modelData = await rootBundle.load('assets/models/model.onnx');
      final options = OrtSessionOptions();
      _session = OrtSession.fromBuffer(modelData.buffer.asUint8List(), options);
    } catch (e) {
      debugPrint('Failed to load VAD model: $e');
    }
  }

  Future<List<(Duration, Duration)>?> detectVoices(
      String filePath, Duration duration) async {
    await loadModel();
    if (_session == null) {
      debugPrint('VAD Session not initialized, cannot detect voices.');
      return null;
    }

    try {
      final options =
          FFmpegPcmConversionOptions(sampleRate: sampleRate, format: 'f32le');

      final result = await _ffmpegProcessingService.processAudio(
        filePath: filePath,
        duration: duration,
        options: options,
        onData: (pcmStream, audioDuration) async {
          final segments = await _performVoiceDetection(
            pcmStream: pcmStream,
            audioDuration: audioDuration,
            options: options,
          );
          return segments;
        },
      );
      return result;
    } catch (e, s) {
      debugPrint('Failed to detect voices in $filePath: $e\n$s');
      return null;
    }
  }

  Future<List<(Duration, Duration)>> _performVoiceDetection({
    required Stream<List<int>> pcmStream,
    required Duration audioDuration,
    required FFmpegPcmConversionOptions options,
  }) async {
    final segments = <(Duration, Duration)>[];
    bool isVoice = false;
    Duration? startTime;

    final zeroData = Float32List.fromList(List.filled(2 * 1 * 128, 0.0));
    var state = OrtValueTensor.createTensorWithDataList(zeroData, [2, 1, 128]);

    final windowedStream = pcmStream.transform(
        WindowingTransformer(windowSize: windowSize, bytesPerSample: 4));

    int samplesProcessed = 0;

    await for (var floatChunk in windowedStream) {
      if (floatChunk.length < windowSize) {
        final paddedChunk = Float32List(windowSize);
        paddedChunk.setAll(0, floatChunk);
        floatChunk = paddedChunk;
      }

      final modelInput = Float32List(windowSize);
      modelInput.setAll(0, floatChunk);

      final inputTensor = OrtValueTensor.createTensorWithDataList(
          modelInput, [1, windowSize]);
      final srTensor = OrtValueTensor.createTensorWithData(options.sampleRate);

      final inputs = {
        'input': inputTensor as OrtValue,
        'state': state,
        'sr': srTensor
      };

      final runOptions = OrtRunOptions();
      final outputs = _session!.run(runOptions, inputs);
      state = outputs[1]! as OrtValueTensor;

      inputTensor.release();
      srTensor.release();

      final scoresTensor = outputs[0] as OrtValueTensor;
      final scores = scoresTensor.value[0] as List<double>;
      final isCurrentChunkVoice = scores[0] > 0.5;

      if (isCurrentChunkVoice && !isVoice) {
        isVoice = true;
        startTime = Duration(
            milliseconds:
                (samplesProcessed / options.sampleRate * 1000).round());
      }

      if (!isCurrentChunkVoice && isVoice) {
        isVoice = false;
        final endTime = Duration(
            milliseconds:
                (samplesProcessed / options.sampleRate * 1000).round());
        if (startTime != null) {
          segments.add((startTime, endTime));
          startTime = null;
        }
      }
      samplesProcessed += windowSize;
    }

    if (isVoice && startTime != null) {
      segments.add((startTime, audioDuration));
    }

    state.release();
    return segments;
  }

  void close() {
    _session?.release();
    _session = null;
  }
}

覆盖libonnxruntime.so

现在flutter插件的onnxruntime,打包的libonnxruntime.so的版本比较低(为1.15),只能支持IR版本为9的onnx格式模型。

而huggingface上这个model.onnx,格式是10。这会导致我们的代码在Arm 64的实机上运行的时候,会因为无法解析onnx模型文件而失败。

为了解决这个问题,我们可以转化onnx模型文件的格式,但是比较复杂。

我想了一个比较简单的方法:

因为flutter中的onnxruntime库,调用的也是libonnxruntime.so这个动态链接库文件,所以我在maven上下载了版本为1.22的microsoft/onnxruntime-android 的libonnxruntime4j_jni.so和libonnxruntime.so,替换了onnxruntime包中自带的,发现可以工作。

替换方法就是把.so文件放入项目的android/app/src/main/jniLibs/arm64-v8a/中,之后使用flutter构建app,让flutter中的onnxruntime库调用这个更新版本的libonnxruntime.so。

Logo

开源鸿蒙跨平台开发社区汇聚开发者与厂商,共建“一次开发,多端部署”的开源生态,致力于降低跨端开发门槛,推动万物智联创新。

更多推荐