File size: 3,177 Bytes
5e23710
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
# 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")