Repository navigation
Expand file tree
/
Copy pathgraph_executor.cpp
More file actions
67 lines (53 loc) · 2.11 KB
/
Copy pathgraph_executor.cpp
File metadata and controls
67 lines (53 loc) · 2.11 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
#include "engine/framework/runtime/graph_executor.h"
#include "engine/framework/core/backend.h"
#include "engine/framework/core/execution_context.h"
#include <algorithm>
#include <chrono>
#include <stdexcept>
namespace engine::runtime {
namespace {
constexpr size_t kDefaultGraphNodeCapacity = 32768;
}
GraphExecutor::GraphExecutor(engine::core::ExecutionContext & execution_context)
: execution_context_(execution_context),
gallocr_(ggml_gallocr_new(
ggml_backend_get_default_buffer_type(execution_context_.backend()))) {}
GraphExecutor::~GraphExecutor() {
if (gallocr_ != nullptr) {
ggml_gallocr_free(gallocr_);
gallocr_ = nullptr;
}
}
GraphRunResult GraphExecutor::run(
ggml_context * ggml_ctx,
ggml_tensor * output,
int warmup,
int iterations) {
if (ggml_ctx == nullptr || output == nullptr) {
throw std::runtime_error("GraphExecutor::run requires non-null graph context and output tensor");
}
if (cached_ggml_ctx_ != ggml_ctx || cached_output_ != output || cached_graph_ == nullptr) {
cached_ggml_ctx_ = ggml_ctx;
cached_output_ = output;
cached_graph_ = ggml_new_graph_custom(ggml_ctx, kDefaultGraphNodeCapacity, false);
ggml_build_forward_expand(cached_graph_, output);
ggml_gallocr_alloc_graph(gallocr_, cached_graph_);
cached_graph_allocated_ = true;
} else if (!cached_graph_allocated_) {
ggml_gallocr_alloc_graph(gallocr_, cached_graph_);
cached_graph_allocated_ = true;
}
for (int i = 0; i < warmup; ++i) {
engine::core::compute_backend_graph(execution_context_.backend(), cached_graph_);
}
const auto started = std::chrono::steady_clock::now();
for (int i = 0; i < std::max(1, iterations); ++i) {
engine::core::compute_backend_graph(execution_context_.backend(), cached_graph_);
}
const auto ended = std::chrono::steady_clock::now();
GraphRunResult result;
result.average_ms =
std::chrono::duration<double, std::milli>(ended - started).count() / std::max(1, iterations);
return result;
}
} // namespace engine::runtime