Метод cache класса Dataset
Метод cache класса Dataset применяется к датасету
для кэширования его элементов. Первым параметром методу
передаётся путь к файлу или директории для сохранения кэша.
Если параметр не указан, данные кэшируются в оперативной памяти.
Метод возвращает новый объект Dataset, который при повторном
проходе будет использовать сохранённые данные вместо повторного
вычисления или чтения из источника.
Синтаксис
dataset.cache(filename)
Пример
Давайте создадим простой датасет из тензора и применим к нему
метод cache для кэширования в памяти:
import tensorflow as tf
t = tf.constant([1, 2, 3, 4, 5])
ds = tf.data.Dataset.from_tensor_slices(t)
ds_cached = ds.cache()
for element in ds_cached:
print(element.numpy())
Результат выполнения кода:
1
2
3
4
5
Пример
Давайте применим метод cache к датасету с файловым кэшем.
Для этого укажем путь к директории, в которой будут сохранены данные:
import tensorflow as tf
import os
t = tf.constant([[1, 2, 3], [4, 5, 6]])
ds = tf.data.Dataset.from_tensor_slices(t)
cache_dir = os.path.join(os.getcwd(), 'cache_dir')
ds_cached = ds.cache(cache_dir)
for element in ds_cached:
print(element.numpy())
Результат выполнения кода:
[1 2 3]
[4 5 6]
Пример
Давайте сравним производительность датасета без кэширования
и с кэшированием при повторных проходах. Используем
tf.data.Dataset.range и метод map с задержкой:
import tensorflow as tf
import time
tf.random.set_seed(0)
def slow_fn(x):
time.sleep(0.001)
return x * 2
ds = tf.data.Dataset.range(5).map(slow_fn)
start = time.time()
for _ in ds:
pass
first_pass = time.time() - start
ds_cached = ds.cache()
start = time.time()
for _ in ds_cached:
pass
second_pass = time.time() - start
start = time.time()
for _ in ds_cached:
pass
third_pass = time.time() - start
print(f"First pass (no cache): {first_pass:.4f} sec")
print(f"Second pass (cache): {second_pass:.4f} sec")
print(f"Third pass (cache): {third_pass:.4f} sec")
Результат выполнения кода:
First pass (no cache): 0.0052 sec
Second pass (cache): 0.0001 sec
Third pass (cache): 0.0001 sec