Метод set_tensor класса Interpreter
Метод set_tensor класса Interpreter устанавливает
значение входного тензора модели TensorFlow Lite. Первым
параметром метод принимает индекс входного тензора, который
можно получить через get_input_details. Вторым параметром
передается значение тензора - обычно это массив NumPy,
содержащий подготовленные входные данные.
Метод не возвращает результат. После установки всех входных
тензоров необходимо вызвать invoke для выполнения
вычислений интерпретатором. Форма и тип данных тензора
должны соответствовать ожиданиям модели, иначе возникнет
ошибка.
Синтаксис
interpreter.set_tensor(tensor_index, value)
Пример
Давайте создадим простой интерпретатор с двумя входными
тензорами и установим их значения через set_tensor,
после чего прочитаем их обратно через get_tensor:
import tensorflow as tf
import numpy as np
interpreter = tf.lite.Interpreter(
model_content=None,
experimental_preserve_all_tensors=True
)
# create a small model-like interpreter manually
interpreter = tf.lite.Interpreter(model_path=None)
Однако для наглядной демонстрации работы метода удобнее
использовать реальную модель TensorFlow Lite. Создадим
простую модель с одним входом и сохраним ее, затем загрузим
в интерпретатор и передадим входные данные через
set_tensor:
import tensorflow as tf
import numpy as np
# build and save a simple model
model = tf.keras.Sequential([
tf.keras.layers.Dense(1, input_shape=(1,))
])
model.save('model.keras')
converter = tf.lite.TFLiteConverter.from_keras_model(model)
tflite_model = converter.convert()
with open('model.tflite', 'wb') as f:
f.write(tflite_model)
# load into interpreter
interpreter = tf.lite.Interpreter(model_path='model.tflite')
interpreter.allocate_tensors()
# get input details
input_details = interpreter.get_input_details()
print(input_details[0]['index'])
Результат выполнения кода:
0
Пример
Давайте установим значение входного тензора и выполним инференс модели:
import tensorflow as tf
import numpy as np
interpreter = tf.lite.Interpreter(model_path='model.tflite')
interpreter.allocate_tensors()
input_details = interpreter.get_input_details()
output_details = interpreter.get_output_details()
input_index = input_details[0]['index']
input_data = np.array([[5.0]], dtype=np.float32)
interpreter.set_tensor(input_index, input_data)
interpreter.invoke()
output_index = output_details[0]['index']
res = interpreter.get_tensor(output_index)
print(res)
Результат выполнения кода:
[[...]]
Пример
Давайте установим несколько входных тензоров, если модель принимает более одного входа, перебирая их индексы в цикле:
Результат выполнения кода:
"done"
Смотрите также
-
класс
Interpreter,
который запускает модели TensorFlow Lite -
метод
get_tensor,
который читает значение тензора интерпретатора -
метод
invoke,
который выполняет вычисления интерпретатора -
метод
get_input_details,
который возвращает информацию о входных тензорах