zhifu gao
2024-03-11 0d9384c8c0161259192cc3d676ca0d60e0d18e5c
funasr/auto/auto_model.py
@@ -155,9 +155,8 @@
            device = "cpu"
            kwargs["batch_size"] = 1
        kwargs["device"] = device
        if kwargs.get("ncpu", 4):
            torch.set_num_threads(kwargs.get("ncpu"))
        torch.set_num_threads(kwargs.get("ncpu", 4))
        
        # build tokenizer
        tokenizer = kwargs.get("tokenizer", None)
@@ -495,11 +494,19 @@
                export_dir = export_utils.export_onnx(
                                        model=model,
                                        data_in=data_list,
                                        quantize=quantize,
                                        fallback_num=fallback_num,
                                        calib_num=calib_num,
                                        opset_version=opset_version,
                                        **kwargs)
            else:
                export_dir = export_utils.export_torchscripts(
                                        model=model,
                                        data_in=data_list,
                                        quantize=quantize,
                                        fallback_num=fallback_num,
                                        calib_num=calib_num,
                                        opset_version=opset_version,
                                        **kwargs)
        return export_dir