Repository navigation
Expand file tree
/
Copy pathgreedy_decode.cpp
More file actions
54 lines (42 loc) · 1.7 KB
/
Copy pathgreedy_decode.cpp
File metadata and controls
54 lines (42 loc) · 1.7 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
#include "engine/framework/sampling/greedy_decode.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 & 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;
}
} // namespace engine::sampling