diff --git a/python/tflite_micro/BUILD b/python/tflite_micro/BUILD index 4780b9a2..25f480c2 100644 --- a/python/tflite_micro/BUILD +++ b/python/tflite_micro/BUILD @@ -94,6 +94,7 @@ py_test( ], deps = [ ":runtime", + requirement("ai-edge-litert"), requirement("numpy"), requirement("tensorflow"), "//tensorflow/lite/micro/examples/recipes:add_four_numbers", diff --git a/python/tflite_micro/runtime_test.py b/python/tflite_micro/runtime_test.py index 2a9003c6..6897ec0d 100644 --- a/python/tflite_micro/runtime_test.py +++ b/python/tflite_micro/runtime_test.py @@ -25,6 +25,7 @@ import weakref import numpy as np import tensorflow as tf +from ai_edge_litert import interpreter as litert_interpreter from tensorflow.python.framework import test_util from tensorflow.python.platform import test from tflite_micro.python.tflite_micro import runtime @@ -199,10 +200,10 @@ class ConvModelTests(test_util.TensorFlowTestCase): tflm_interpreter = runtime.Interpreter.from_bytes(model_data) # TFLite interpreter - tflite_interpreter = tf.lite.Interpreter( + tflite_interpreter = litert_interpreter.Interpreter( model_content=model_data, experimental_op_resolver_type=\ - tf.lite.experimental.OpResolverType.BUILTIN_REF) + litert_interpreter.OpResolverType.BUILTIN_REF) tflite_interpreter.allocate_tensors() tflite_output_details = tflite_interpreter.get_output_details()[0] tflite_input_details = tflite_interpreter.get_input_details()[0] diff --git a/tensorflow/lite/micro/examples/hello_world/BUILD b/tensorflow/lite/micro/examples/hello_world/BUILD index 988b7dd6..d327dc35 100644 --- a/tensorflow/lite/micro/examples/hello_world/BUILD +++ b/tensorflow/lite/micro/examples/hello_world/BUILD @@ -53,6 +53,7 @@ py_binary( "@absl_py//absl:app", "@absl_py//absl/flags", "@absl_py//absl/logging", + requirement("ai-edge-litert"), requirement("numpy"), requirement("tensorflow"), "//python/tflite_micro:runtime", diff --git a/tensorflow/lite/micro/examples/hello_world/evaluate.py b/tensorflow/lite/micro/examples/hello_world/evaluate.py index 8b6f9488..f6f7ac1b 100644 --- a/tensorflow/lite/micro/examples/hello_world/evaluate.py +++ b/tensorflow/lite/micro/examples/hello_world/evaluate.py @@ -16,6 +16,7 @@ import os import tensorflow as tf from absl import app from absl import flags +from ai_edge_litert import interpreter as litert_interpreter import numpy as np import matplotlib.pyplot as plt from tensorflow.python.platform import resource_loader @@ -92,9 +93,9 @@ def get_tflm_prediction(model_path, x_values): # returns the prediction of the interpreter. def get_tflite_prediction(model_path, x_values): # TFLite interpreter - tflite_interpreter = tf.lite.Interpreter( + tflite_interpreter = litert_interpreter.Interpreter( model_path=model_path, - experimental_op_resolver_type=tf.lite.experimental.OpResolverType. + experimental_op_resolver_type=litert_interpreter.OpResolverType. BUILTIN_REF, ) tflite_interpreter.allocate_tensors() diff --git a/tensorflow/lite/micro/examples/mnist_lstm/BUILD b/tensorflow/lite/micro/examples/mnist_lstm/BUILD index 7d818b21..e8c5d2b6 100644 --- a/tensorflow/lite/micro/examples/mnist_lstm/BUILD +++ b/tensorflow/lite/micro/examples/mnist_lstm/BUILD @@ -6,6 +6,7 @@ py_binary( srcs = ["train.py"], srcs_version = "PY3", deps = [ + requirement("ai-edge-litert"), requirement("numpy"), requirement("tensorflow"), ], diff --git a/tensorflow/lite/micro/examples/mnist_lstm/evaluate_test.py b/tensorflow/lite/micro/examples/mnist_lstm/evaluate_test.py index a7d74cd3..f5f0ad53 100644 --- a/tensorflow/lite/micro/examples/mnist_lstm/evaluate_test.py +++ b/tensorflow/lite/micro/examples/mnist_lstm/evaluate_test.py @@ -17,6 +17,7 @@ import os import numpy as np import tensorflow as tf +from ai_edge_litert import interpreter as litert_interpreter from tensorflow.python.framework import test_util from tensorflow.python.platform import resource_loader from tensorflow.python.platform import test @@ -43,10 +44,10 @@ class LSTMFloatModelTest(test_util.TensorFlowTestCase): evaluate.predict_image(self.tflm_interpreter, wrong_size_image_path) def testCompareWithTFLite(self): - tflite_interpreter = tf.lite.Interpreter( + tflite_interpreter = litert_interpreter.Interpreter( model_path=self.model_path, experimental_op_resolver_type=\ - tf.lite.experimental.OpResolverType.BUILTIN_REF) + litert_interpreter.OpResolverType.BUILTIN_REF) tflite_interpreter.allocate_tensors() tflite_output_details = tflite_interpreter.get_output_details()[0] tflite_input_details = tflite_interpreter.get_input_details()[0] diff --git a/tensorflow/lite/micro/tools/BUILD b/tensorflow/lite/micro/tools/BUILD index 2d1e1874..4d1976e0 100644 --- a/tensorflow/lite/micro/tools/BUILD +++ b/tensorflow/lite/micro/tools/BUILD @@ -34,6 +34,7 @@ py_library( srcs_version = "PY3", visibility = ["//:__subpackages__"], deps = [ + requirement("ai-edge-litert"), "//tensorflow/lite/python:schema_py", ], ) @@ -208,6 +209,7 @@ py_binary( ":model_transforms_utils", "@absl_py//absl:app", "@absl_py//absl/flags", + requirement("ai-edge-litert"), requirement("tensorflow"), "//python/tflite_micro:runtime", "//tensorflow/lite/tools:flatbuffer_utils", diff --git a/tensorflow/lite/micro/tools/generate_test_for_model.py b/tensorflow/lite/micro/tools/generate_test_for_model.py index 8c5b4070..7bc77e95 100644 --- a/tensorflow/lite/micro/tools/generate_test_for_model.py +++ b/tensorflow/lite/micro/tools/generate_test_for_model.py @@ -18,6 +18,7 @@ import csv import numpy as np import tensorflow as tf +from ai_edge_litert import interpreter as litert_interpreter from tflite_micro.tensorflow.lite.python import schema_py_generated as schema_fb @@ -103,9 +104,9 @@ class TestDataGenerator: if (len(self.model_paths) != 1): raise RuntimeError(f'Single model expected') model_path = self.model_paths[0] - interpreter = tf.lite.Interpreter(model_path=model_path, + interpreter = litert_interpreter.Interpreter(model_path=model_path, experimental_op_resolver_type=\ - tf.lite.experimental.OpResolverType.BUILTIN_REF) + litert_interpreter.OpResolverType.BUILTIN_REF) interpreter.allocate_tensors() @@ -140,10 +141,10 @@ class TestDataGenerator: for model_path in self.model_paths: # Load model and run a single inference with random inputs. - interpreter = tf.lite.Interpreter( + interpreter = litert_interpreter.Interpreter( model_path=model_path, experimental_op_resolver_type=\ - tf.lite.experimental.OpResolverType.BUILTIN_REF) + litert_interpreter.OpResolverType.BUILTIN_REF) interpreter.allocate_tensors() input_tensor = interpreter.tensor( interpreter.get_input_details()[0]['index']) diff --git a/tensorflow/lite/micro/tools/layer_by_layer_debugger.py b/tensorflow/lite/micro/tools/layer_by_layer_debugger.py index 8aa263b0..441f523f 100644 --- a/tensorflow/lite/micro/tools/layer_by_layer_debugger.py +++ b/tensorflow/lite/micro/tools/layer_by_layer_debugger.py @@ -20,6 +20,7 @@ import unittest from absl import app from absl import flags from absl import logging +from ai_edge_litert import interpreter as litert_interpreter import numpy as np import tensorflow as tf @@ -194,7 +195,7 @@ def main(_) -> None: intrepreter_config=runtime.InterpreterConfig.kPreserveAllTensors, ) - tflite_interpreter = tf.lite.Interpreter( + tflite_interpreter = litert_interpreter.Interpreter( model_path=_INPUT_TFLITE_FILE.value, experimental_preserve_all_tensors=True, ) diff --git a/third_party/python_requirements.in b/third_party/python_requirements.in index 29c081e5..d5a076c6 100644 --- a/third_party/python_requirements.in +++ b/third_party/python_requirements.in @@ -34,3 +34,4 @@ mako pillow yapf protobuf +ai-edge-litert diff --git a/third_party/python_requirements.txt b/third_party/python_requirements.txt index b0d91331..2aded403 100644 --- a/third_party/python_requirements.txt +++ b/third_party/python_requirements.txt @@ -11,6 +11,9 @@ absl-py==2.0.0 \ # keras # tensorboard # tensorflow +ai-edge-litert==1.0.1 \ + --hash=sha256:25a9b1577941498842bf77630722eda1163026c37abd57af66791a6955551b9d + # via -r third_party/python_requirements.in astunparse==1.6.3 \ --hash=sha256:5ad93a8456f0d084c3456d059fd9a92cce667963232cbf763eac3bc5b7940872 \ --hash=sha256:c2652417f2c8b5bb325c885ae329bdf3f86424075c4fd1a128674bc6fba4b8e8 @@ -505,6 +508,7 @@ numpy==1.26.3 \ --hash=sha256:f73497e8c38295aaa4741bdfa4fda1a5aedda5473074369eca10626835445511 # via # -r third_party/python_requirements.in + # ai-edge-litert # h5py # keras # ml-dtypes