返回

Project SonicCraft Qinling:将 Spleeter 转换为 TFLite

SonicCraft Qinling 需要在端侧完成人声与伴奏的分离。Spleeter 在听感上尚可,但原生 TensorFlow 模型体积偏大,难以直接部署到目标设备,因此将其导出并压缩为 TensorFlow Lite。

实际耗时主要不在最后的 convert 调用,而在 Spleeter 提供的权重往往并非一份可直接转换的 SavedModel。

环境

pip install tensorflow tensorflow-model-optimization

同时安装 Spleeter。转换脚本与 Spleeter 宜置于同一 Python 环境;分处不同环境时,容易在 import 与版本依赖上反复排查。

获取预训练权重

本地没有权重则无法导出。可先执行一次分离,由工具自动下载模型:

spleeter separate -i input_audio_file.wav -p spleeter:2stems -o output

权重通常位于 ~/.spleeter/pretrained_models,具体路径随安装方式略有差异。此处使用 2stems(人声与伴奏);多轨模型更大,现阶段不必引入。

导出 SavedModel

Spleeter 内部多为 checkpoint 一类格式,而 TFLite Converter 对 SavedModel 支持更好,故先完成一次导出:

import tensorflow as tf
from spleeter.separator import Separator

separator = Separator("spleeter:2stems")
export_dir = "./spleeter_saved_model"

tf.saved_model.save(separator.model, export_dir)

该步骤在不同 Spleeter 版本上表现并不一致:有的版本可直接 save,有的则因签名不完整而失败。导出出错时,宜先核对版本组合,再调整转换参数。

转换为 tflite

import tensorflow as tf

saved_model_dir = "./spleeter_saved_model"
converter = tf.lite.TFLiteConverter.from_saved_model(saved_model_dir)

converter.optimizations = [tf.lite.Optimize.DEFAULT]

tflite_model = converter.convert()

with open("spleeter_model.tflite", "wb") as f:
    f.write(tflite_model)

convert 失败,常见原因包括:计算图中存在 TFLite 尚不支持的算子、动态 shape,或上一阶段 SavedModel 本身不完整。宜先保证转换成功,再考虑进一步量化。

量化

.tflite 体积仍不理想,可启用量化。动态范围量化成本最低,上述 Optimize.DEFAULT 往往已包含类似优化:

converter.optimizations = [tf.lite.Optimize.DEFAULT]

整数量化需要提供校准数据,且输入 shape 须与模型一致:

def representative_dataset():
    for _ in range(num_calibration_steps):
        yield [input_data]

converter.optimizations = [tf.lite.Optimize.DEFAULT]
converter.representative_dataset = representative_dataset

分离质量最终仍取决于听感。体积下降而听感明显劣化,则压缩失去意义。

加载与验证

生成文件并不等于可在端侧运行。可先在本机用 Interpreter 做一次推理:

import tensorflow as tf

interpreter = tf.lite.Interpreter(model_path="spleeter_model.tflite")
interpreter.allocate_tensors()

input_details = interpreter.get_input_details()
output_details = interpreter.get_output_details()

interpreter.set_tensor(input_details[0]["index"], input_data)
interpreter.invoke()

output_data = interpreter.get_tensor(output_details[0]["index"])

input_data 应沿用训练阶段的预处理。shape 不匹配时,set_tensor 会立即报错,便于定位问题。


后续若接入 OpenHarmony 或约束更严的 runtime,瓶颈更可能出现在算子支持与输入管线,而非重复编写导出脚本。