trevdatastreams's picture
Publish deterministic Arm NN parser-table OOB PoC
277dba7 verified
Raw
History Blame Contribute Delete
4.63 kB
#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;
}
}