blob: 9c754deb34f43141ea1a6c1e3b266cf465740c6f [file] [edit]
/*
* Copyright 2024 The ChromiumOS Authors
* Use of this source code is governed by a BSD-style license that can be
* found in the LICENSE file.
*/
#include <gtest/gtest.h>
#include <iostream>
#include "absl/random/random.h"
#include "common/async_driver.h"
#include "common/simple_model_builder.h"
#include "tensorflow/lite/core/c/builtin_op_data.h"
#include "tensorflow/lite/tools/command_line_flags.h"
#include "tensorflow/lite/tools/delegates/delegate_provider.h"
#include "tensorflow/lite/tools/tool_params.h"
namespace tflite::cros::tests {
using DelegateProvider = tools::DelegateProvider;
using ProvidedDelegateList = tools::ProvidedDelegateList;
using ProvidedDelegate = tools::ProvidedDelegateList::ProvidedDelegate;
using ToolParams = tools::ToolParams;
using TfLiteDelegatePtr = Interpreter::TfLiteDelegatePtr;
class Environment : public ::testing::Environment {
public:
Environment() : delegate_list_(&params_) {}
~Environment() override {}
bool InitFromCommandLine(int* argc, char** argv) {
constexpr char kSettingsFlagName[] = "stable_delegate_settings_file";
std::string settings;
// Use our own flags directly instead of delegate_list_.AppendCmdlineFlags()
// to make the settings flag required.
std::vector<tflite::Flag> flags = {
Flag::CreateFlag(kSettingsFlagName, &settings,
"The path to the delegate settings JSON file.",
Flag::kRequired),
};
if (!tflite::Flags::Parse(argc, const_cast<const char**>(argv), flags)) {
std::cout << Flags::Usage(argv[0], flags) << std::endl;
return false;
}
delegate_list_.AddAllDelegateParams();
params_.Set<std::string>(kSettingsFlagName, settings);
return true;
}
TfLiteDelegatePtr CreateDelegate() {
std::vector<ProvidedDelegate> provided_delegates =
delegate_list_.CreateAllRankedDelegates();
EXPECT_EQ(provided_delegates.size(), 1);
if (provided_delegates.empty()) {
return tools::CreateNullDelegate();
}
return std::move(provided_delegates[0].delegate);
}
void SetUp() override {
// Ensure that the delegate can be created before running any test.
ASSERT_NE(CreateDelegate(), nullptr);
}
private:
ToolParams params_;
ProvidedDelegateList delegate_list_;
};
Environment* g_env = nullptr;
TEST(AsyncDelegate, Inference) {
using TensorArgs = SimpleModelBuilder::TensorArgs;
const TensorArgs base_args = {
.type = kTfLiteFloat32,
.shape = {2, 3},
};
auto arg_with_name = [&](const char* name) {
TensorArgs args = base_args;
args.name = name;
return args;
};
SimpleModelBuilder mb;
int a = mb.AddInput(arg_with_name("a"));
int b = mb.AddInput(arg_with_name("b"));
int c = mb.AddInput(arg_with_name("c"));
int d = mb.AddOutput(arg_with_name("d"));
int a_plus_b = mb.AddInternalTensor(arg_with_name("a_plus_b"));
mb.AddOperator<TfLiteAddParams>({
.op = kTfLiteBuiltinAdd,
.inputs = {a, b},
.outputs = {a_plus_b},
});
mb.AddOperator<TfLiteSubParams>({
.op = kTfLiteBuiltinSub,
.inputs = {a_plus_b, c},
.outputs = {d},
});
auto model = mb.Build();
ASSERT_NE(model, nullptr);
TfLiteDelegatePtr delegate = g_env->CreateDelegate();
ASSERT_NE(delegate, nullptr);
auto driver = AsyncDriver::Create(std::move(delegate), std::move(model));
ASSERT_NE(driver, nullptr);
ASSERT_EQ(driver->Prepare(), kTfLiteOk);
int n = 1;
for (int dim : base_args.shape) {
n *= dim;
}
absl::BitGen gen;
std::vector<float> a_data(n);
std::vector<float> b_data(n);
std::vector<float> c_data(n);
std::vector<float> d_data(n);
for (int i = 0; i < n; ++i) {
a_data[i] = absl::Uniform(gen, 0.0, 1.0);
b_data[i] = absl::Uniform(gen, 0.0, 1.0);
c_data[i] = absl::Uniform(gen, 0.0, 1.0);
d_data[i] = a_data[i] + b_data[i] - c_data[i];
}
EXPECT_EQ(driver->SetInputTensor("a", a_data), kTfLiteOk);
EXPECT_EQ(driver->SetInputTensor("b", b_data), kTfLiteOk);
EXPECT_EQ(driver->SetInputTensor("c", c_data), kTfLiteOk);
ASSERT_EQ(driver->Invoke(), kTfLiteOk);
auto output = driver->GetOutputTensor<float>("d");
for (int i = 0; i < n; ++i) {
EXPECT_FLOAT_EQ(output[i], d_data[i]);
}
}
} // namespace tflite::cros::tests
using tflite::cros::tests::g_env;
int main(int argc, char** argv) {
g_env = new tflite::cros::tests::Environment();
if (!g_env->InitFromCommandLine(&argc, argv)) {
delete g_env;
return EXIT_FAILURE;
}
testing::InitGoogleTest(&argc, argv);
// GoogleTest takes the ownership of g_env.
::testing::AddGlobalTestEnvironment(g_env);
return RUN_ALL_TESTS();
}