File size: 3,583 Bytes
819e690
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
// RNNoise AX SDK:用 AX Engine 替换原版 compute_rnn(网络推理),
// 其余信号处理(biquad/FFT/pitch/合成)沿用原版 C 实现。
#include "rnnoise_ax.hpp"

#include "model_runner.hpp"

#include <cstring>
#include <mutex>
#include <stdexcept>
#include <vector>

extern "C" {
#include "denoise.h"
#include "rnn.h"
}

namespace {

// 单实例模型 runner(compute_rnn 无状态上下文可用,采用全局注册方式;
// 同一时刻只允许一个 RNNoiseAX 实例执行推理)。
std::mutex g_runner_mu;
ModelRunner* g_ax_runner = nullptr;

extern "C" void rnnoise_ax_set_runner(ModelRunner* runner) {
    std::lock_guard<std::mutex> lk(g_runner_mu);
    g_ax_runner = runner;
}

extern "C" void compute_rnn(const RNNoise* model, RNNState* rnn,
                            float* gains, float* vad,
                            const float* input, int arch) {
    (void)model;
    (void)arch;
    ModelRunner* runner = nullptr;
    {
        std::lock_guard<std::mutex> lk(g_runner_mu);
        runner = g_ax_runner;
    }
    if (runner == nullptr) {
        throw std::runtime_error("compute_rnn: AX runner 未初始化");
    }

    std::vector<std::vector<float>> feeds = {
        std::vector<float>(input, input + NB_FEATURES),
        std::vector<float>(rnn->conv1_state,
                           rnn->conv1_state + CONV1_STATE_SIZE),
        std::vector<float>(rnn->conv2_state,
                           rnn->conv2_state + CONV2_STATE_SIZE),
        std::vector<float>(rnn->gru1_state,
                           rnn->gru1_state + GRU1_OUT_SIZE),
        std::vector<float>(rnn->gru2_state,
                           rnn->gru2_state + GRU2_OUT_SIZE),
        std::vector<float>(rnn->gru3_state,
                           rnn->gru3_state + GRU3_OUT_SIZE),
    };
    std::vector<std::vector<float>> outs = runner->Run(feeds);

    std::memcpy(gains, outs[0].data(), NB_BANDS * sizeof(float));
    *vad = outs[1][0];
    std::memcpy(rnn->conv1_state, outs[2].data(),
                CONV1_STATE_SIZE * sizeof(float));
    std::memcpy(rnn->conv2_state, outs[3].data(),
                CONV2_STATE_SIZE * sizeof(float));
    std::memcpy(rnn->gru1_state, outs[4].data(),
                GRU1_OUT_SIZE * sizeof(float));
    std::memcpy(rnn->gru2_state, outs[5].data(),
                GRU2_OUT_SIZE * sizeof(float));
    std::memcpy(rnn->gru3_state, outs[6].data(),
                GRU3_OUT_SIZE * sizeof(float));
}

}  // namespace

struct RNNoiseAX::Impl {
    std::unique_ptr<ModelRunner> runner;
    DenoiseState* state = nullptr;

    explicit Impl(const std::string& model_path)
        : runner(new ModelRunner(model_path)),
          state(rnnoise_create(nullptr)) {
        if (state == nullptr) {
            throw std::runtime_error("rnnoise_create 失败");
        }
        rnnoise_ax_set_runner(runner.get());
    }

    ~Impl() {
        rnnoise_ax_set_runner(nullptr);
        if (state) {
            rnnoise_destroy(state);
        }
    }
};

RNNoiseAX::RNNoiseAX(const std::string& model_path)
    : impl_(new Impl(model_path)) {}

RNNoiseAX::~RNNoiseAX() = default;

void RNNoiseAX::Reset() {
    if (!impl_) {
        return;
    }
    rnnoise_ax_set_runner(nullptr);
    rnnoise_destroy(impl_->state);
    impl_->state = rnnoise_create(nullptr);
    if (impl_->state == nullptr) {
        throw std::runtime_error("rnnoise_create 失败");
    }
    rnnoise_ax_set_runner(impl_->runner.get());
}

float RNNoiseAX::ProcessFrame(float* out, const float* in) {
    return rnnoise_process_frame(impl_->state, out, in);
}