Класс FixedLengthRecordDataset
Класс FixedLengthRecordDataset создает датасет из
одного или нескольких файлов, содержимое которых
разбивается на записи фиксированной длины.
Первым параметром передаются имена файлов,
вторым - длина одной записи в байтах.
Третьим параметром можно передать длину заголовка,
который пропускается в начале каждого файла,
четвертым - длину хвоста, который пропускается в конце.
Такой формат часто используется для бинарных
наборов данных, например CIFAR-10.
Синтаксис
tf.data.FixedLengthRecordDataset(
filenames, record_bytes, header_bytes=0, footer_bytes=0, buffer_size=0
)
Пример
Давайте создадим бинарный файл, в котором
каждая запись занимает ровно 4 байта,
и прочитаем его как датасет:
import tensorflow as tf
import os
path = "records.bin"
with open(path, "wb") as f:
f.write(b"1234")
f.write(b"5678")
f.write(b"90ab")
dataset = tf.data.FixedLengthRecordDataset([path], record_bytes=4)
for record in dataset:
print(record)
os.remove(path)
Результат выполнения кода:
tf.Tensor(b'1234', shape=(), dtype=string)
tf.Tensor(b'5678', shape=(), dtype=string)
tf.Tensor(b'90ab', shape=(), dtype=string)
Пример
Давайте пропустим заголовок размером 2 байта
в начале файла и хвост размером 2 байта в конце:
import tensorflow as tf
import os
path = "records.bin"
with open(path, "wb") as f:
f.write(b"HH")
f.write(b"1234")
f.write(b"5678")
f.write(b"TT")
dataset = tf.data.FixedLengthRecordDataset(
[path], record_bytes=4, header_bytes=2, footer_bytes=2
)
for record in dataset:
print(record)
os.remove(path)
Результат выполнения кода:
tf.Tensor(b'1234', shape=(), dtype=string)
tf.Tensor(b'5678', shape=(), dtype=string)
Пример
Давайте преобразуем байтовые записи в числа
с помощью метода map и выведем результат:
import tensorflow as tf
import os
path = "records.bin"
with open(path, "wb") as f:
f.write(bytes([1, 2, 3, 4]))
f.write(bytes([5, 6, 7, 8]))
dataset = tf.data.FixedLengthRecordDataset([path], record_bytes=4)
dataset = dataset.map(lambda x: tf.io.decode_raw(x, tf.uint8))
for record in dataset:
print(record)
os.remove(path)
Результат выполнения кода:
tf.Tensor([1 2 3 4], shape=(4,), dtype=uint8)
tf.Tensor([5 6 7 8], shape=(4,), dtype=uint8)
Смотрите также
-
класс
TFRecordDataset,
который читает датасет из файлов формата TFRecord -
класс
TextLineDataset,
который читает датасет построчно из текстовых файлов -
класс
Iterator,
который перебирает элементы датасета -
функцию
split_dataset,
которая разбивает датасет на части