Метод scan класса Dataset
Метод scan применяется к объекту класса Dataset
и выполняет последовательное преобразование его элементов.
В отличие от метода map, который обрабатывает
каждый элемент независимо, scan передаёт
накопленное состояние из предыдущего шага в следующий.
Первым параметром передаётся функция преобразования,
принимающая текущее состояние и очередной элемент датасета.
Вторым параметром передаётся начальное состояние.
Синтаксис
dataset.scan(initial_state, scan_func)
Пример
Давайте создадим датасет из чисел
1, 2, 3, 4, 5
и последовательно накопим их сумму:
import tensorflow as tf
ds = tf.data.Dataset.from_tensor_slices([1, 2, 3, 4, 5])
def scan_func(state, element):
return state + element, state + element
scanned = ds.scan(0, scan_func)
for res in scanned.as_numpy_iterator():
print(res)
Результат выполнения кода:
1
3
6
10
15
Пример
Давайте применим scan к датасету и вернём
в качестве результата кортеж из состояния и элемента:
import tensorflow as tf
ds = tf.data.Dataset.from_tensor_slices([1, 2, 3, 4, 5])
def scan_func(state, element):
new_state = state + element
return new_state, (new_state, element)
scanned = ds.scan(0, scan_func)
for res in scanned.as_numpy_iterator():
print(res)
Результат выполнения кода:
(1, 1)
(3, 2)
(6, 3)
(10, 4)
(15, 5)
Пример
Давайте используем scan для последовательного
умножения элементов датасета:
Результат выполнения кода:
1
2
6
24
120