载入中...
搜索中...
未找到
OnnxModel.cpp
浏览该文件的文档.
3
4#include <set>
5
6namespace eve::tensor {
10OnnxModel::OnnxModel(std::unique_ptr<Impl> impl) : impl_(std::move(impl)) {}
11OnnxModel::~OnnxModel() = default;
12OnnxModelInfo OnnxModel::info() const { return impl_->model.info; }
14 try {
15 auto impl = std::make_unique<Impl>();
17 return Result<std::unique_ptr<OnnxModel>>::success(std::unique_ptr<OnnxModel>(new OnnxModel(std::move(impl))));
18 } catch (const onnx_detail::Failure& e) {
19 return Result<std::unique_ptr<OnnxModel>>::failure(
20 Diagnostic::error(e.code, e.what(), {}, {}, "tensor.onnx.import"));
21 } catch (const std::exception& e) {
22 return Result<std::unique_ptr<OnnxModel>>::failure(
23 Diagnostic::error(DiagnosticCode::Failed, e.what(), {}, {}, "tensor.onnx.import"));
24 }
25}
26Result<std::vector<OnnxNamedTensor>> OnnxModel::run(std::span<const OnnxNamedTensor> feeds,
27 std::span<const std::string> requested,
28 OnnxRunOptions options) const {
29 return runInternal(feeds, requested, nullptr, options);
30}
31Result<OnnxGpuResult> OnnxModel::runGpu(std::span<const OnnxNamedTensor> feeds, OnnxCompute& compute,
32 std::span<const std::string> requested, OnnxRunOptions options) const {
33 class Counter final : public OnnxCompute {
34 public:
36 size_t count = 0;
37 explicit Counter(OnnxCompute& t) : target(t) {}
38 Result<std::vector<uint8_t>> dispatch(const OnnxKernel& k) override {
39 auto r = target.dispatch(k);
40 if (r.ok()) ++count;
41 return r;
42 }
43 Result<OnnxBuffer> enqueue(const std::string& source, std::span<const OnnxBuffer> inputs, size_t bytes,
44 uint32_t work) override {
45 auto r = target.enqueue(source, inputs, bytes, work);
46 if (r.ok()) ++count;
47 return r;
48 }
49 } counter(compute);
50 compute.beginRun();
51 struct RunScope {
53 ~RunScope() { device.endRun(); }
54 } scope{compute};
55 auto r = runInternal(feeds, requested, &counter, options);
56 if (!r.ok()) return Result<OnnxGpuResult>::failure(r.status());
57 return Result<OnnxGpuResult>::success({std::move(r.value()), counter.count, compute.transferStats()});
58}
59Result<std::vector<OnnxNamedTensor>> OnnxModel::runInternal(std::span<const OnnxNamedTensor> feeds,
60 std::span<const std::string> requested,
62 using namespace onnx_detail;
63 std::string nodeName;
64 try {
65 const auto& model = impl_->model;
66 const std::vector<std::string> targets =
67 requested.empty() ? model.info.outputs : std::vector<std::string>(requested.begin(), requested.end());
68 std::set<std::string> seenFeeds;
69 for (const auto& f : feeds) {
70 nodeName = f.name;
71 if (count(f.tensor.shape) * elementSize(f.tensor.element) != f.tensor.bytes.size())
72 throw Failure("Tensor byte count mismatch");
73 auto input = model.inputs.find(f.name);
74 if (input == model.inputs.end() || !seenFeeds.insert(f.name).second)
75 throw Failure("Unknown or duplicate feed");
76 if (static_cast<int>(f.tensor.element) != input->second.type ||
77 f.tensor.shape.size() != input->second.shape.size())
78 throw Failure("Feed type/rank mismatch");
79 for (size_t i = 0; i < f.tensor.shape.size(); ++i)
80 if (input->second.shape[i] >= 0 && input->second.shape[i] != f.tensor.shape[i])
81 throw Failure("Feed dimension mismatch");
82 }
83 nodeName.clear();
84 auto result = evaluate(model, feeds, targets, compute, options);
85 return Result<std::vector<OnnxNamedTensor>>::success(std::move(result));
86 } catch (const Failure& e) {
87 return Result<std::vector<OnnxNamedTensor>>::failure(
88 Diagnostic::error(e.code, e.what(), e.path.empty() ? nodeName : e.path, {}, "tensor.onnx.run"));
89 } catch (const std::exception& e) {
90 return Result<std::vector<OnnxNamedTensor>>::failure(
91 Diagnostic::error(DiagnosticCode::Failed, e.what(), nodeName, {}, "tensor.onnx.run"));
92 }
93}
94} // namespace eve::tensor
LogicalId target
void * impl
EvpackChunkInput input
Definition Evpack.cpp:170
vkb::Device & device
int inputs
Definition GridGraph.cpp:23
double r
JobScope scope
std::uint64_t bytes
OnnxCompute * compute
float f
std::string path
Definition PlayHost.cpp:110
float t
glm::mat4 model
std::uint32_t count
const SquirrelValueOptions & options
const UnitySourceAsset & source
static Diagnostic error(DiagnosticCode code, std::string message, std::string path={}, DiagnosticDetails details={}, std::string source={})
Construct an error diagnostic with the standard error severity.
Definition Diagnostic.h:125
Move-only operation result carrying either a value or Status.
Definition Result.h:155
static Result success(T value)
Construct a successful result owning value.
Definition Result.h:164
static Result failure(Status status)
Construct a failed result from a structured status.
Definition Result.h:175
GPU execution boundary for native ONNX; retains no model and retains compiled resources for the lifet...
Definition OnnxCompute.h:25
virtual void endRun() noexcept
Retire temporary activations/recordings; compiled resources may survive for later calls.
Definition OnnxCompute.h:71
Native ONNX import and CPU/GPU execution using tensor kernels, without ONNX Runtime.
Definition OnnxModel.h:63
Result< std::vector< OnnxNamedTensor > > run(std::span< const OnnxNamedTensor > feeds, std::span< const std::string > requested={}, OnnxRunOptions options={}) const
Execute selected named values, or graph outputs when requested is empty.
Definition OnnxModel.cpp:26
Result< OnnxGpuResult > runGpu(std::span< const OnnxNamedTensor > feeds, OnnxCompute &compute, std::span< const std::string > requested={}, OnnxRunOptions options={}) const
Execute with GPU matrix, convolution, normalization, activation and resampling kernels.
Definition OnnxModel.cpp:31
OnnxModelInfo info() const
Return owning graph diagnostics; importing does not imply full operator support.
Definition OnnxModel.cpp:12
static Result< std::unique_ptr< OnnxModel > > load(std::span< const uint8_t > bytes)
Import bounded in-memory ONNX ModelProto; copies all retained data.
Definition OnnxModel.cpp:13
~OnnxModel()
Destroy model-owned graph and packed initializers; outputs remain valid.
ModelData parse(std::span< const uint8_t > bytes)
Parse.
Synchronous GPU kernel request; all inputs are borrowed only until dispatch returns.
Definition OnnxCompute.h:12
Owning admission report; unsupported nodes remain inspectable but cannot execute.
Definition OnnxModel.h:48
onnx_detail::ModelData model
Definition OnnxModel.cpp:8
Per-call deterministic RNG and optional strict finite-output diagnostic.
Definition OnnxModel.h:37