MCPcopy Create free account
hub / github.com/ARM-software/armnn / Run

Method Run

tests/InferenceModel.hpp:544–608  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

542 }
543
544 std::chrono::duration<double, std::milli> Run(
545 const std::vector<armnnUtils::TContainer>& inputContainers,
546 std::vector<armnnUtils::TContainer>& outputContainers)
547 {
548 for (unsigned int i = 0; i < outputContainers.size(); ++i)
549 {
550 const unsigned int expectedOutputDataSize = GetOutputSize(i);
551
552 mapbox::util::apply_visitor([expectedOutputDataSize, i](auto&& value)
553 {
554 const unsigned int actualOutputDataSize = armnn::numeric_cast<unsigned int>(value.size());
555 if (actualOutputDataSize < expectedOutputDataSize)
556 {
557 unsigned int outputIndex = i;
558 throw armnn::Exception(
559 fmt::format("Not enough data for output #{0}: expected "
560 "{1} elements, got {2}", outputIndex, expectedOutputDataSize, actualOutputDataSize));
561 }
562 },
563 outputContainers[i]);
564 }
565
566 std::shared_ptr<armnn::IProfiler> profiler = m_Runtime->GetProfiler(m_NetworkIdentifier);
567
568 // Start timer to record inference time in EnqueueWorkload (in milliseconds)
569 const auto start_time = armnn::GetTimeNow();
570
571 armnn::Status ret;
572 if (m_ImportInputsIfAligned)
573 {
574 std::vector<armnn::ImportedInputId> importedInputIds = m_Runtime->ImportInputs(
575 m_NetworkIdentifier, MakeInputTensors(inputContainers), armnn::MemorySource::Malloc);
576
577 std::vector<armnn::ImportedOutputId> importedOutputIds = m_Runtime->ImportOutputs(
578 m_NetworkIdentifier, MakeOutputTensors(outputContainers), armnn::MemorySource::Malloc);
579
580 ret = m_Runtime->EnqueueWorkload(m_NetworkIdentifier,
581 MakeInputTensors(inputContainers),
582 MakeOutputTensors(outputContainers),
583 importedInputIds,
584 importedOutputIds);
585 }
586 else
587 {
588 ret = m_Runtime->EnqueueWorkload(m_NetworkIdentifier,
589 MakeInputTensors(inputContainers),
590 MakeOutputTensors(outputContainers));
591 }
592 const auto duration = armnn::GetTimeDuration(start_time);
593
594 // if profiling is enabled print out the results
595 if (profiler && profiler->IsProfilingEnabled())
596 {
597 profiler->Print(std::cout);
598 }
599
600 if (ret == armnn::Status::Failure)
601 {

Callers 1

InferenceTestFunction · 0.45

Calls 13

formatEnum · 0.85
GetTimeNowFunction · 0.85
GetTimeDurationFunction · 0.85
ExceptionClass · 0.50
MakeInputTensorsFunction · 0.50
MakeOutputTensorsFunction · 0.50
sizeMethod · 0.45
GetProfilerMethod · 0.45
ImportInputsMethod · 0.45
ImportOutputsMethod · 0.45
EnqueueWorkloadMethod · 0.45
IsProfilingEnabledMethod · 0.45

Tested by 1

InferenceTestFunction · 0.36