DiariZen / model_convert /export_cnn_onnx.py
inoryQwQ's picture
Upload folder using huggingface_hub
37c4768 verified
Raw
History Blame
3.45 kB
#!/usr/bin/env python3
from pathlib import Path
import json, sys, struct
import toml, numpy as np, torch, torchaudio
import onnx, onnxruntime as ort
TASK_DIR = Path(__file__).resolve().parents[1]
EXPORT_DIR = Path(__file__).resolve().parent
ORIGIN = TASK_DIR / "origin"
_hf = Path((TASK_DIR / "cache/acquire/hf_model_path.txt").read_text().strip())
HF_PATH = _hf if _hf.is_absolute() else (TASK_DIR.parents[2] / _hf).resolve()
sys.path.insert(0, str(ORIGIN))
sys.path.insert(0, str(ORIGIN / "pyannote-audio"))
from diarizen.models.eend.model_wavlm_conformer import Model
def main():
torch.manual_seed(0)
out_dir = EXPORT_DIR
out_dir.mkdir(parents=True, exist_ok=True)
config = toml.load(HF_PATH / "config.toml")
args = config["model"]["args"]
model = Model(**args)
state = torch.load(HF_PATH / "pytorch_model.bin", map_location="cpu")
model.load_state_dict(state, strict=True)
model.eval()
wav_path = ORIGIN / "example/EN2002a_30s.wav"
waveform, sr = torchaudio.load(wav_path)
if sr != 16000:
waveform = torchaudio.functional.resample(waveform, sr, 16000)
waveform = waveform[:1, :4*16000]
full_onnx = onnx.load(TASK_DIR / "export_4s/model.onnx")
eps = 1e-5
for n in full_onnx.graph.node:
if n.name == '/Constant_2':
eps = struct.unpack('f', n.attribute[0].t.raw_data)[0]
wav_1d = waveform[0, :4*16000]
mean = wav_1d.mean()
var = ((wav_1d - mean) ** 2).mean()
wav_ln = (wav_1d - mean) / np.sqrt(var + eps)
wav_input = torch.from_numpy(wav_ln).float().unsqueeze(0)
wavlm = model.wavlm_model
with torch.no_grad():
cnn_features, _ = wavlm.feature_extractor(wav_input, None)
print(f"CNN features shape: {cnn_features.shape}")
class CNNWrapper(torch.nn.Module):
def __init__(self, fe):
super().__init__()
self.fe = fe
def forward(self, x):
out, _ = self.fe(x, None)
return out
cnn_wrapper = CNNWrapper(wavlm.feature_extractor)
cnn_wrapper.eval()
onnx_path = out_dir / "cnn_features.onnx"
torch.onnx.export(
cnn_wrapper, wav_input, onnx_path.as_posix(),
input_names=["waveform_ln"], output_names=["cnn_features"],
opset_version=16, do_constant_folding=True, dynamic_axes=None,
)
onnx_model = onnx.load(onnx_path.as_posix())
onnx.checker.check_model(onnx_model)
sess = ort.InferenceSession(onnx_path.as_posix(), providers=["CPUExecutionProvider"])
ort_out = sess.run(None, {"waveform_ln": wav_input.numpy()})[0]
diff = ort_out - cnn_features.numpy()
cosine = float(np.dot(ort_out.ravel(), cnn_features.numpy().ravel()) /
(np.linalg.norm(ort_out.ravel()) * np.linalg.norm(cnn_features.numpy().ravel()) + 1e-12))
print(f"Cosine: {cosine:.10f}, MAE: {float(np.mean(np.abs(diff))):.8f}")
calib_dir = out_dir / "calib_data"
calib_dir.mkdir(exist_ok=True)
np.save(calib_dir / "waveform_0.npy", wav_input.numpy().astype(np.float32))
meta = {
"model_name": "diarizen-wavlm-cnn-frontend",
"inputs": [{"name": "waveform_ln", "shape": [1, 64000], "dtype": "float32"}],
"outputs": [{"name": "cnn_features", "shape": list(ort_out.shape), "dtype": "float32"}],
}
(out_dir / "model_meta.json").write_text(json.dumps(meta, indent=2) + "\n")
print(f"Done: {onnx_path}")
if __name__ == "__main__":
main()