虎牙科技
虎牙公司作为国内领先的电竞直播平台,一直秉承技术驱动内容发展的理念,在大流量业务视频数据处理、 5G超高清视频直播技术,以及实时内容创作和直播互动技术等关键领域持续突破,致力于降低直播内容生产门槛,优化产业生态环境,为直播行业发展赋能。
WeNet 应用于虎牙直播的消音项目中,所谓的消音指的是实时屏蔽直播过程出现的黑词语音;要达到这一目的,必须准确且快速地对实时语音流进行识别。虽然WeNet原有的 Libtorch 方案可以满足这一要求,但为了降低 RTF 及节省更多的计算资源,我们采用了 ONNX 的推理方案,实验表明,在流式识别模式下,ONNX 比 Libtorch相对提速20%左右,且 WER 基本一致。
目前我们已经把 ONNX 的推理方案开源到了 WeNet 的官方 repo 中,开源地址详见:
https://github.com/wenet-e2e/wenet/blob/main/runtime/core/decoder/onnx_asr_model.h
下文我们先直接介绍最终的实验结果,再介绍 ONNX 推理实现的一些技术细节。
实验 ONNX vs libtorch
为了对比libtorch和ONNX的推理性能,我们在一台60核的机器上做了多线程解码测试,结果如下表所示。可以看到,ONNX 在不同并发情况,不同的解码参数设置下,相对 libtorch 都有一定程度的加速。但我们也注意到,随着线程数增加,ONNX 的的加速效果逐渐下降。

runtime中ONNX推理实现
ONNX的推理流程为:加载模型和超参数、初始化cache、encoder推理、ctc推理,最后是rescore推理。onnx_asr_model和torch_asr_model都继承自asr_model,asr_model中定义了Reset、ForwardEncoderFunc、AttentionRescoring这三个虚函数,Reset实现了offset_、att_cache_等cache的初始化;ForwardEncoderFunc包含了encoder和ctc推理;AttentionRescoring对识别结果做最后的重打分。总体而言,导出onnx后,参照onnx的官方示例、api文档及libtorch的推理流程,实现runtime的onnx推理并不难,但具体实现过程还是出现了一些问题,以下是一些技术细节问题和我们的解决方案。
ONNX 线程数的配置
刚开始开发时没有把onnx的线程数设置1,导致出现了比libtorch提速百分之五六十的情况,后来才发现onnx默认用了多核去加速解码。设置onnx线程数的代码为:
session_options_.SetIntraOpNumThreads(num_threads);
session_options_.SetInterOpNumThreads(num_threads); 优雅读取asr网络的超参数
刚开始把runtime需要用到的一些网络超参数写成一个txt文件,然后runtime里一行行地读取赋值,明显这种方式不够优雅,也容易出错;后来发现可以把这些参数以一个字典的格式附带保存到onnx模型里,这种方式与libtorch的超参保存方式类似,不容易出错,也更优雅。读取超参数的代码为:
auto model_metadata = encoder_session_->GetModelMetadata();
Ort::AllocatorWithDefaultOptions allocator;
encoder_output_size_ = std::move(
atoi(model_metadata.LookupCustomMetadataMap("output_size", allocator))); encoder的入参个数不定
由于导出onnx时,不同的chunk_size和num_decoding_left_chunks组合配置,encoder.onnx的入参是不一样的(onnx会优化掉无用参数),如果通过手工的方式统计不同组合时,encoder.onnx分别需要哪些输入,那代码逻辑将变得晦涩复杂;因此,对于encoder,会先获取其输入参数名列表,然后在准备encoder的输入vector时,会根据参数名列表,挑选相应变量加入vector。而对于encoder的输出、ctc和decoder的输入输出,也全部采用直接从模型读取参数名列表的方式,避免手工定义参数名列表。
//根据encoder_in_names_准备输入vector
std::vector<Ort::Value> inputs;
for (auto name : encoder_in_names_) {
if (!strcmp(name, "chunk")) {
inputs.emplace_back(std::move(feats_ort));
} else if (!strcmp(name, "offset")) {
inputs.emplace_back(std::move(offset_ort));
} else if (!strcmp(name, "required_cache_size")) {
inputs.emplace_back(std::move(required_cache_size_ort));
} else if (!strcmp(name, "att_cache")) {
inputs.emplace_back(std::move(att_cache_ort_));
} else if (!strcmp(name, "cnn_cache")) {
inputs.emplace_back(std::move(cnn_cache_ort_));
} else if (!strcmp(name, "att_mask")) {
inputs.emplace_back(std::move(att_mask_ort));
}
} int类型参数
在python中,encoder的两个输入参数offset和required_cache_size都是int类型,因此在runtime时,这两个参数的构造比较特殊,CreateTensor函数里shape和shape_len两个形参应分别传入空指针和0。此处感谢@Mddct大佬给出的解决方法。
Ort::Value offset_ort = Ort::Value::CreateTensor<int64_t>(
memory_info, &offset_int64, 1, std::vector<int64_t>{}.data(), 0); 变量的全局与局部设置
由于encoder的两个入参att_cache_ort_和cnn_cache_ort_存放的是历史缓存,且Reset函数也需要初始化这两个变量,因此必须设置成全局变量。而att_mask_ort需要设置成局部变量主要有三个原因:
att_mask_ort需根据offset_动态设置里面的值; 构造encoder的输入vector时会通过std::move把att_mask_ort清空; Reset函数不需对att_mask_ort进行初始化。
用于构造att_cache_ort_的att_cache_这个vector也必须设置成全局变量,因为onnxruntime库没有对att_cache_里的数据进行拷贝,而只是维持了指向att_cache_的指针。如果在Reset函数中将att_cache_声明为一个局部变量,再拿去构造att_cache_ort_,在识别wav.scp时,会出现跑着跑着突然崩溃的情况。原因是att_cache_是一个局部变量,内存已经被系统回收。cnn_cache_需设置成全局也是同理。Ort变量env_也需设置成全局,原因是env_ 维持着其他对象使用的日志记录状态,必须在使用onnxruntime的其它函数之前创建好env_。此处感谢@Mddct 和@Duum 两位大佬指出bug并给出解释。
虎牙多巴胺团队
虎牙多巴胺团队专注智能内容安全能力构建,团队中有语音,NLP,视觉等多个领域的专家。未来,虎牙多巴胺团队也会持续助力开源,助力WeNet。欢迎各位对我们团队感兴趣,有志于语音的同学加入虎牙多巴胺团队,详情联系huanghuiyan@huya.com。
