
往期回顾
一:ONNX 不支持 torch.tensor 转 index 的切片操作
二:不支持传入 NoneType 类型参数
三:不支持 List[tensor] 形式的输入和输出
四:支持 If-else 的动态变化1. 引入 torch.jit.script 和 slice_helper
2. 引入音频长度为 1、值为 0 的 cache 张量替代 None
3. 通过 cat 将 list 中的 tensor 合并为一个输出
4. 引入 torch.jit.script,将 if-else 部分独立成函数 get_next_cache_start在4中使用的 torch.jit.Script 需要将要修饰的代码抽离出本体;此外,被torch.jit.Script 修饰的代码会导致模型在 Pytorch 上无法进行计算,即无法进行训练。为了保证训练的正常进行,原方案不得不引入大量的 if onnx_mode 分支以区分导出和训练。 在2中引入的长度为 1 的张量使得推理过程和训练过程不匹配。这是由于推理过程对输入和输出的 slice 需要根据 cache 长度计算,因此这里的长度 1 会导致推理和训练的 slice 不一致,需要做特殊处理。在原方案中,该特殊处理即引入大量的 onnx 分支以区分导出和训练。
ifonnx_mode: # 原方案onnx_mode分支示例导出分支else:训练分支优化策略
1. opset >= 13 时 ONNX 已经可以支持上述的切片操作(链接 https://github.com/onnx/onnx/blob/main/docs/Operators.md#Slice 由 @Mddct 提供)
2. 引入音频长度为 0、值为 0 的 dummy 张量替代 None
3. 在我们最新版的 forward_chunk 接口中(https://github.com/wenet-e2e/wenet/pull/1002),在设计时就考虑到了该情况,输入输出已经不再是 List 形式
4. 通过巧妙地构造输入的 cache 形状,保证了整个 inference 流程中遇到的所有 if-else 在给定的解码策略下永远都只走同一个分支,也即 if-else 从动态变为静态,这时 ONNX 导出时报的类似如下的 warning 便可以愉快忽略!/home/xcsong/workspace/wenet/wenet/transformer/encoder.py:226: TracerWarning: Converting a tensor to a Python boolean might cause the trace to be incorrect. We can't record the data flow of Python values, so this value will be treated as a constant in the future. This means that the trace might not generalize to other inputs!if required_cache_size < 0:a = torch.ones((1, 2, 0, 4))b = torch.ones((1, 2, 3, 4))c = torch.cat((a, b), dim=2)torch.equal(b, c)# Trued = torch.split(a, 2, dim=-1)torch.equal(d[0], d[1])# Trueifrequired_cache_size <0:next_cache_start =0elifrequired_cache_size ==0:next_cache_start = attention_key_sizeelse:next_cache_start = max(attention_key_size - required_cache_size,0)补充介绍
五:最初的方案只涉及 U2 模型,对于 U2++ 模型中双向 decoder 的导出,反向 decoder 需要构造反向输入,其中涉及到 pad_sequence 这个 op,ONNX 是不支持导出的。
六:超参数的存取问题,不希望单独写到一个文件中,最好可以和 ONNX 模型耦合我们的解决方案:
5. 针对该问题,由@Mddct重新设计了一个与 pad_sequence 等价且能被 ONNX 感知到 shape 变化的函数 https://github.com/wenet-e2e/wenet/blob/main/wenet/transformer/asr_model.py#L683-L721# NOTE(Mddct): `pad_sequence` is not supported by ONNX, it is used# in `reverse_pad_list` thus we have to refine the below code.# Issue: https://github.com/wenet-e2e/wenet/issues/1113# Equal to:# >>> r_hyps = reverse_pad_list(r_hyps, r_hyps_lens, float(self.ignore_id))# >>> r_hyps, _ = add_sos_eos(r_hyps, self.sos, self.eos, self.ignore_id)max_len = torch.max(r_hyps_lens)index_range = torch.arange(0, max_len,1).to(encoder_out.device)seq_len_expand = r_hyps_lens.unsqueeze(1)seq_mask = seq_len_expand > index_range# (beam, max_len)# >>> seq_mask# >>> tensor([[ True, True, True],# >>> [ True, True, True],# >>> [ True, False, False]])index = (seq_len_expand -1) - index_range# (beam, max_len)# >>> index# >>> tensor([[ 2, 1, 0],# >>> [ 2, 1, 0],# >>> [ 0, -1, -2]])index = index * seq_mask# >>> index# >>> tensor([[2, 1, 0],# >>> [2, 1, 0],# >>> [0, 0, 0]])r_hyps = torch.gather(r_hyps,1, index)# >>> r_hyps# >>> tensor([[3, 2, 1],# >>> [4, 8, 9],# >>> [2, 2, 2]])r_hyps = torch.where(seq_mask, r_hyps, self.eos)# >>> r_hyps# >>> tensor([[3, 2, 1],# >>> [4, 8, 9],# >>> [2, eos, eos]])r_hyps = torch.cat([hyps[:,0:1], r_hyps], dim=1)# >>> r_hyps# >>> tensor([[sos, 3, 2, 1],# >>> [sos, 4, 8, 9],# >>> [sos, 2, eos, eos]])6. 通过 onnx 的 metadata 接口将超参数全部存入 onnx 模型# writeonnx_encoder = onnx.load(encoder_outpath)for(k, v)inargs.items():meta = onnx_encoder.metadata_props.add()meta.key, meta.value = str(k), str(v)onnx.save(onnx_encoder, encoder_outpath)# readort_session = onnxruntime.InferenceSession(encoder_outpath)meta = ort_session.get_modelmeta()print("\t\tcustom_metadata_map={}".format(meta.custom_metadata_map))