| #include <armnnDeserializer/IDeserializer.hpp> |
| #include <armnn/IRuntime.hpp> |
|
|
| #include <cstdint> |
| #include <fstream> |
| #include <iostream> |
| #include <iterator> |
| #include <string> |
| #include <vector> |
|
|
| void PrintOutputBinding( |
| const armnnDeserializer::IDeserializer& deserializer, |
| const std::string& name) |
| { |
| const armnnDeserializer::BindingPointInfo binding = |
| deserializer.GetNetworkOutputBindingInfo(0, name); |
| const armnn::TensorShape& shape = binding.m_TensorInfo.GetShape(); |
|
|
| std::cout << name << " binding=" << binding.m_BindingId << " shape=["; |
| for (unsigned int i = 0; i < shape.GetNumDimensions(); ++i) |
| { |
| if (i != 0) |
| { |
| std::cout << ","; |
| } |
| std::cout << shape[i]; |
| } |
| std::cout << "] elements=" << binding.m_TensorInfo.GetNumElements() << "\n"; |
| } |
|
|
| void PrintRuntimeOutput( |
| const armnn::IRuntime& runtime, |
| armnn::NetworkId networkId, |
| armnn::LayerBindingId bindingId) |
| { |
| const armnn::TensorInfo& info = |
| runtime.GetOutputTensorInfo(networkId, bindingId); |
| const armnn::TensorShape& shape = info.GetShape(); |
|
|
| std::cout << "runtime output " << bindingId << " shape=["; |
| for (unsigned int i = 0; i < shape.GetNumDimensions(); ++i) |
| { |
| if (i != 0) |
| { |
| std::cout << ","; |
| } |
| std::cout << shape[i]; |
| } |
| std::cout << "] elements=" << info.GetNumElements() << "\n"; |
| } |
|
|
| int main(int argc, char** argv) |
| { |
| if (argc != 2) |
| { |
| std::cerr << "usage: " << argv[0] << " MODEL.armnn\n"; |
| return 2; |
| } |
|
|
| std::ifstream input(argv[1], std::ios::binary); |
| if (!input) |
| { |
| std::cerr << "could not open " << argv[1] << "\n"; |
| return 2; |
| } |
|
|
| std::vector<uint8_t> bytes( |
| (std::istreambuf_iterator<char>(input)), |
| std::istreambuf_iterator<char>()); |
|
|
| try |
| { |
| auto deserializer = armnnDeserializer::IDeserializer::Create(); |
| auto network = deserializer->CreateNetworkFromBinary(bytes); |
| std::cout << "loaded " << bytes.size() << " bytes\n"; |
| PrintOutputBinding(*deserializer, "output-0"); |
| PrintOutputBinding(*deserializer, "output-1"); |
|
|
| auto runtime = armnn::IRuntime::Create(armnn::IRuntime::CreationOptions()); |
| auto optimized = armnn::Optimize( |
| *network, {armnn::Compute::CpuRef}, runtime->GetDeviceSpec()); |
| if (!optimized) |
| { |
| std::cerr << "optimization failed\n"; |
| return 1; |
| } |
|
|
| armnn::NetworkId networkId = -1; |
| std::string errorMessage; |
| if (runtime->LoadNetwork(networkId, std::move(optimized), errorMessage) |
| != armnn::Status::Success) |
| { |
| std::cerr << "runtime load failed: " << errorMessage << "\n"; |
| return 1; |
| } |
|
|
| PrintRuntimeOutput(*runtime, networkId, 0); |
| PrintRuntimeOutput(*runtime, networkId, 1); |
|
|
| auto inputBinding = |
| deserializer->GetNetworkInputBindingInfo(0, "input"); |
| auto output0Binding = |
| deserializer->GetNetworkOutputBindingInfo(0, "output-0"); |
| auto output1Binding = |
| deserializer->GetNetworkOutputBindingInfo(0, "output-1"); |
|
|
| inputBinding.m_TensorInfo.SetConstant(true); |
| std::vector<float> inputData{1.0F, 2.0F, 3.0F, 4.0F}; |
| std::vector<float> output0Data( |
| output0Binding.m_TensorInfo.GetNumElements()); |
| std::vector<float> output1Data( |
| output1Binding.m_TensorInfo.GetNumElements()); |
|
|
| armnn::InputTensors inputTensors{{ |
| inputBinding.m_BindingId, |
| armnn::ConstTensor(inputBinding.m_TensorInfo, inputData.data()) |
| }}; |
| armnn::OutputTensors outputTensors{ |
| { |
| output0Binding.m_BindingId, |
| armnn::Tensor(output0Binding.m_TensorInfo, output0Data.data()) |
| }, |
| { |
| output1Binding.m_BindingId, |
| armnn::Tensor(output1Binding.m_TensorInfo, output1Data.data()) |
| } |
| }; |
|
|
| const armnn::Status enqueueStatus = |
| runtime->EnqueueWorkload(networkId, inputTensors, outputTensors); |
| std::cout << "enqueue status=" |
| << (enqueueStatus == armnn::Status::Success ? "success" : "failure") |
| << "\n"; |
| std::cout << "allocated output elements=" |
| << output0Data.size() << "," << output1Data.size() << "\n"; |
| return network ? 0 : 1; |
| } |
| catch (const std::exception& error) |
| { |
| std::cerr << "load failed: " << error.what() << "\n"; |
| return 1; |
| } |
| } |
|
|