Repository navigation
Expand file tree
/
Copy pathactivity_graph.cpp
More file actions
55 lines (42 loc) · 1.51 KB
/
Copy pathactivity_graph.cpp
File metadata and controls
55 lines (42 loc) · 1.51 KB
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
#include "engine/framework/audio/activity_graph.h"
#include <stdexcept>
namespace engine::sampling {
namespace {
const engine::core::ModulePortSpec kSingleInput[] = {
{"input", engine::core::PortKind::Activation, false},
};
const engine::core::ModulePortSpec kSingleOutput[] = {
{"output", engine::core::PortKind::Activation, false},
};
const engine::core::ModuleSchema kVADGateSchema = {
"VADGate",
"sampling.gating",
kSingleInput,
1,
kSingleOutput,
1,
"Thresholds frame energy into a binary speech mask.",
};
}
VADGateModule::VADGateModule(VADGateConfig config) : config_(config) {}
const VADGateConfig & VADGateModule::config() const noexcept {
return config_;
}
const engine::core::ModuleSchema & VADGateModule::schema() const noexcept {
return static_schema();
}
engine::core::TensorValue VADGateModule::build(engine::core::ModuleBuildContext & ctx, const engine::core::TensorValue & energy) const {
if (ctx.ggml == nullptr) {
throw std::runtime_error("ModuleBuildContext.ggml is null");
}
engine::core::validate_rank_between(energy, 2, 2, "energy");
auto shifted = engine::core::wrap_tensor(
ggml_scale_bias(ctx.ggml, energy.tensor, 1.0f, -config_.threshold),
energy.shape,
GGML_TYPE_F32);
return engine::core::wrap_tensor(ggml_step(ctx.ggml, shifted.tensor), energy.shape, GGML_TYPE_F32);
}
const engine::core::ModuleSchema & VADGateModule::static_schema() noexcept {
return kVADGateSchema;
}
} // namespace engine::sampling