Repository navigation
Expand file tree
/
Copy pathdecode_modules.cpp
More file actions
90 lines (72 loc) · 2.83 KB
/
Copy pathdecode_modules.cpp
File metadata and controls
90 lines (72 loc) · 2.83 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
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
#include "engine/framework/sampling/decode_modules.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 kGreedyDecodeSchema = {
"GreedyDecode",
"sampling.decode",
kSingleInput,
1,
kSingleOutput,
1,
"Takes argmax over the vocabulary axis for each decoding step.",
};
const engine::core::ModuleSchema kVADGateSchema = {
"VADGate",
"sampling.gating",
kSingleInput,
1,
kSingleOutput,
1,
"Thresholds frame energy into a binary speech mask.",
};
}
const engine::core::ModuleSchema & GreedyDecodeModule::schema() const noexcept {
return static_schema();
}
engine::core::TensorValue GreedyDecodeModule::build(engine::core::ModuleBuildContext & ctx, const engine::core::TensorValue & logits) const {
if (ctx.ggml == nullptr) {
throw std::runtime_error("ModuleBuildContext.ggml is null");
}
engine::core::validate_rank_between(logits, 3, 3, "logits");
auto flat = engine::core::reshape_tensor(
ctx,
engine::core::ensure_backend_addressable_layout(ctx, logits),
engine::core::TensorShape::from_dims({logits.shape.num_elements() / logits.shape.last_dim(), logits.shape.last_dim()}));
auto argmax = engine::core::wrap_tensor(
ggml_argmax(ctx.ggml, flat.tensor),
engine::core::TensorShape::from_dims({flat.shape.dims[0]}),
GGML_TYPE_I32);
return engine::core::reshape_tensor(ctx, argmax, engine::core::TensorShape::from_dims({logits.shape.dims[0], logits.shape.dims[1]}));
}
const engine::core::ModuleSchema & GreedyDecodeModule::static_schema() noexcept {
return kGreedyDecodeSchema;
}
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