GCRN / python /gcrn_sdk /inference.py
inoryQwQ's picture
Upload folder using huggingface_hub
5e23710 verified
Raw
History Blame Contribute Delete
3.18 kB
# coding: utf-8
"""GCRN 语音增强推理:4 秒 chunk STFT -> axmodel -> ISTFT -> 拼接输出。"""
from __future__ import annotations
import json
import time
from pathlib import Path
import numpy as np
from .audio import (
SAMPLE_RATE,
chunks,
istft,
read_wav,
stft,
write_wav,
)
from .runtime import InferenceSession, first_output
INPUT_NAME = "stft_input"
OUTPUT_NAME = "stft_output"
INPUT_SHAPE = (1, 2, 401, 161)
class GCRNDenoiser:
"""AX650 GCRN 语音增强器。默认 .axmodel + axengine。"""
def __init__(self, model_path: str | Path, backend: str | None = None):
self.session = InferenceSession(model_path, backend)
if len(self.session.inputs) != 1 or self.session.inputs[0].shape != INPUT_SHAPE:
raise RuntimeError(f"unexpected GCRN input metadata: {self.session.inputs}")
if len(self.session.outputs) != 1 or self.session.outputs[0].shape != INPUT_SHAPE:
raise RuntimeError(f"unexpected GCRN output metadata: {self.session.outputs}")
self.input_name = self.session.inputs[0].name
def _warmup(self) -> None:
feed = stft(np.zeros(64000, dtype=np.float32))
first_output(self.session, self.session.run({self.input_name: feed}))
def enhance_chunk(self, chunk: np.ndarray) -> np.ndarray:
feed = stft(chunk)
outputs = self.session.run({self.input_name: feed})
output = np.asarray(first_output(self.session, outputs), dtype=np.float32)
return istft(output)
def enhance(self, pcm: np.ndarray) -> np.ndarray:
"""int16 PCM -> 增强后 float32 波形(与输入等长)。"""
self._warmup()
enhanced = [
self.enhance_chunk(chunk)
for chunk, _valid in chunks(pcm)
]
return np.concatenate(enhanced)[: pcm.size]
def enhance_file(
self, input_path: str | Path, output_path: str | Path
) -> dict:
"""输入 16kHz 单声道 WAV,输出增强 WAV + 性能报告 dict。"""
samples, sample_rate = read_wav(input_path)
if sample_rate != SAMPLE_RATE:
raise ValueError(
f"GCRN requires {SAMPLE_RATE} Hz audio, got {sample_rate}"
)
start = time.perf_counter()
enhanced = self.enhance(samples)
elapsed = time.perf_counter() - start
write_wav(output_path, enhanced, sample_rate)
audio_seconds = samples.size / sample_rate
report = {
"backend": self.session.backend,
"model": str(self.session.path),
"input": str(input_path),
"output": str(output_path),
"input_samples": int(samples.size),
"output_samples": int(enhanced.size),
"elapsed_seconds": round(elapsed, 4),
"audio_seconds": round(audio_seconds, 4),
"real_time_factor": round(elapsed / audio_seconds, 4),
}
return report
def save_report(self, report: dict, path: str | Path) -> None:
path = Path(path)
path.parent.mkdir(parents=True, exist_ok=True)
path.write_text(json.dumps(report, indent=2) + "\n", encoding="utf-8")