tf.data:構建 TensorFlow 輸入流水線

在 TensorFlow.org 上檢視 在 Google Colab 中執行 在 GitHub 上檢視原始碼 下載筆記本

tf.data API 使您能夠透過簡單、可重用的元件構建複雜的輸入流水線。例如,影像模型的流水線可以從分散式檔案系統中的檔案中聚合資料,對每張影像應用隨機擾動,並將隨機選擇的影像合併為一個批次進行訓練。文字模型的流水線可能涉及從原始文字資料中提取符號,使用查詢表將它們轉換為嵌入識別符號,並將不同長度的序列批處理在一起。tf.data API 使得處理大量資料、從不同資料格式中讀取以及執行復雜的轉換成為可能。

tf.data API 引入了 tf.data.Dataset 抽象,它表示一個元素序列,其中每個元素由一個或多個分量組成。例如,在影像流水線中,一個元素可以是一個單獨的訓練示例,其中包含一對錶示影像及其標籤的張量分量。

建立資料集有兩種不同的方式

  • 資料源 (source) 根據儲存在記憶體中或一個或多個檔案中的資料構建 Dataset

  • 資料轉換 (transformation) 根據一個或多個 tf.data.Dataset 物件構建資料集。

import tensorflow as tf
2024-08-15 01:37:36.963860: E external/local_xla/xla/stream_executor/cuda/cuda_fft.cc:485] Unable to register cuFFT factory: Attempting to register factory for plugin cuFFT when one has already been registered
2024-08-15 01:37:36.985171: E external/local_xla/xla/stream_executor/cuda/cuda_dnn.cc:8454] Unable to register cuDNN factory: Attempting to register factory for plugin cuDNN when one has already been registered
2024-08-15 01:37:36.991452: E external/local_xla/xla/stream_executor/cuda/cuda_blas.cc:1452] Unable to register cuBLAS factory: Attempting to register factory for plugin cuBLAS when one has already been registered
import pathlib
import os
import matplotlib.pyplot as plt
import pandas as pd
import numpy as np

np.set_printoptions(precision=4)

基本機制

要建立輸入流水線,必須從資料開始。例如,要從記憶體中的資料構建 Dataset,可以使用 tf.data.Dataset.from_tensors()tf.data.Dataset.from_tensor_slices()。或者,如果您的輸入資料儲存在推薦的 TFRecord 格式的檔案中,可以使用 tf.data.TFRecordDataset()

一旦擁有了一個 Dataset 物件,就可以透過在 tf.data.Dataset 物件上鍊式呼叫方法將其轉換為新的 Dataset。例如,可以應用諸如 Dataset.map 的逐元素轉換,以及諸如 Dataset.batch 的多元素轉換。有關轉換的完整列表,請參閱 tf.data.Dataset 的文件。

Dataset 物件是一個 Python 可迭代物件。這使得可以使用 for 迴圈來消耗其元素

dataset = tf.data.Dataset.from_tensor_slices([8, 3, 0, 8, 2, 1])
dataset
WARNING: All log messages before absl::InitializeLog() is called are written to STDERR
I0000 00:00:1723685859.835217   44933 cuda_executor.cc:1015] successful NUMA node read from SysFS had negative value (-1), but there must be at least one NUMA node, so returning NUMA node zero. See more at https://github.com/torvalds/linux/blob/v6.0/Documentation/ABI/testing/sysfs-bus-pci#L344-L355
I0000 00:00:1723685859.839003   44933 cuda_executor.cc:1015] successful NUMA node read from SysFS had negative value (-1), but there must be at least one NUMA node, so returning NUMA node zero. See more at https://github.com/torvalds/linux/blob/v6.0/Documentation/ABI/testing/sysfs-bus-pci#L344-L355
I0000 00:00:1723685859.842691   44933 cuda_executor.cc:1015] successful NUMA node read from SysFS had negative value (-1), but there must be at least one NUMA node, so returning NUMA node zero. See more at https://github.com/torvalds/linux/blob/v6.0/Documentation/ABI/testing/sysfs-bus-pci#L344-L355
I0000 00:00:1723685859.846561   44933 cuda_executor.cc:1015] successful NUMA node read from SysFS had negative value (-1), but there must be at least one NUMA node, so returning NUMA node zero. See more at https://github.com/torvalds/linux/blob/v6.0/Documentation/ABI/testing/sysfs-bus-pci#L344-L355
I0000 00:00:1723685859.858030   44933 cuda_executor.cc:1015] successful NUMA node read from SysFS had negative value (-1), but there must be at least one NUMA node, so returning NUMA node zero. See more at https://github.com/torvalds/linux/blob/v6.0/Documentation/ABI/testing/sysfs-bus-pci#L344-L355
I0000 00:00:1723685859.861635   44933 cuda_executor.cc:1015] successful NUMA node read from SysFS had negative value (-1), but there must be at least one NUMA node, so returning NUMA node zero. See more at https://github.com/torvalds/linux/blob/v6.0/Documentation/ABI/testing/sysfs-bus-pci#L344-L355
I0000 00:00:1723685859.865105   44933 cuda_executor.cc:1015] successful NUMA node read from SysFS had negative value (-1), but there must be at least one NUMA node, so returning NUMA node zero. See more at https://github.com/torvalds/linux/blob/v6.0/Documentation/ABI/testing/sysfs-bus-pci#L344-L355
I0000 00:00:1723685859.868512   44933 cuda_executor.cc:1015] successful NUMA node read from SysFS had negative value (-1), but there must be at least one NUMA node, so returning NUMA node zero. See more at https://github.com/torvalds/linux/blob/v6.0/Documentation/ABI/testing/sysfs-bus-pci#L344-L355
I0000 00:00:1723685859.871403   44933 cuda_executor.cc:1015] successful NUMA node read from SysFS had negative value (-1), but there must be at least one NUMA node, so returning NUMA node zero. See more at https://github.com/torvalds/linux/blob/v6.0/Documentation/ABI/testing/sysfs-bus-pci#L344-L355
I0000 00:00:1723685859.874859   44933 cuda_executor.cc:1015] successful NUMA node read from SysFS had negative value (-1), but there must be at least one NUMA node, so returning NUMA node zero. See more at https://github.com/torvalds/linux/blob/v6.0/Documentation/ABI/testing/sysfs-bus-pci#L344-L355
I0000 00:00:1723685859.878307   44933 cuda_executor.cc:1015] successful NUMA node read from SysFS had negative value (-1), but there must be at least one NUMA node, so returning NUMA node zero. See more at https://github.com/torvalds/linux/blob/v6.0/Documentation/ABI/testing/sysfs-bus-pci#L344-L355
I0000 00:00:1723685859.881840   44933 cuda_executor.cc:1015] successful NUMA node read from SysFS had negative value (-1), but there must be at least one NUMA node, so returning NUMA node zero. See more at https://github.com/torvalds/linux/blob/v6.0/Documentation/ABI/testing/sysfs-bus-pci#L344-L355
I0000 00:00:1723685861.098140   44933 cuda_executor.cc:1015] successful NUMA node read from SysFS had negative value (-1), but there must be at least one NUMA node, so returning NUMA node zero. See more at https://github.com/torvalds/linux/blob/v6.0/Documentation/ABI/testing/sysfs-bus-pci#L344-L355
I0000 00:00:1723685861.100277   44933 cuda_executor.cc:1015] successful NUMA node read from SysFS had negative value (-1), but there must be at least one NUMA node, so returning NUMA node zero. See more at https://github.com/torvalds/linux/blob/v6.0/Documentation/ABI/testing/sysfs-bus-pci#L344-L355
I0000 00:00:1723685861.102280   44933 cuda_executor.cc:1015] successful NUMA node read from SysFS had negative value (-1), but there must be at least one NUMA node, so returning NUMA node zero. See more at https://github.com/torvalds/linux/blob/v6.0/Documentation/ABI/testing/sysfs-bus-pci#L344-L355
I0000 00:00:1723685861.104281   44933 cuda_executor.cc:1015] successful NUMA node read from SysFS had negative value (-1), but there must be at least one NUMA node, so returning NUMA node zero. See more at https://github.com/torvalds/linux/blob/v6.0/Documentation/ABI/testing/sysfs-bus-pci#L344-L355
I0000 00:00:1723685861.106309   44933 cuda_executor.cc:1015] successful NUMA node read from SysFS had negative value (-1), but there must be at least one NUMA node, so returning NUMA node zero. See more at https://github.com/torvalds/linux/blob/v6.0/Documentation/ABI/testing/sysfs-bus-pci#L344-L355
I0000 00:00:1723685861.108307   44933 cuda_executor.cc:1015] successful NUMA node read from SysFS had negative value (-1), but there must be at least one NUMA node, so returning NUMA node zero. See more at https://github.com/torvalds/linux/blob/v6.0/Documentation/ABI/testing/sysfs-bus-pci#L344-L355
I0000 00:00:1723685861.110218   44933 cuda_executor.cc:1015] successful NUMA node read from SysFS had negative value (-1), but there must be at least one NUMA node, so returning NUMA node zero. See more at https://github.com/torvalds/linux/blob/v6.0/Documentation/ABI/testing/sysfs-bus-pci#L344-L355
I0000 00:00:1723685861.112117   44933 cuda_executor.cc:1015] successful NUMA node read from SysFS had negative value (-1), but there must be at least one NUMA node, so returning NUMA node zero. See more at https://github.com/torvalds/linux/blob/v6.0/Documentation/ABI/testing/sysfs-bus-pci#L344-L355
I0000 00:00:1723685861.114046   44933 cuda_executor.cc:1015] successful NUMA node read from SysFS had negative value (-1), but there must be at least one NUMA node, so returning NUMA node zero. See more at https://github.com/torvalds/linux/blob/v6.0/Documentation/ABI/testing/sysfs-bus-pci#L344-L355
I0000 00:00:1723685861.116014   44933 cuda_executor.cc:1015] successful NUMA node read from SysFS had negative value (-1), but there must be at least one NUMA node, so returning NUMA node zero. See more at https://github.com/torvalds/linux/blob/v6.0/Documentation/ABI/testing/sysfs-bus-pci#L344-L355
I0000 00:00:1723685861.117904   44933 cuda_executor.cc:1015] successful NUMA node read from SysFS had negative value (-1), but there must be at least one NUMA node, so returning NUMA node zero. See more at https://github.com/torvalds/linux/blob/v6.0/Documentation/ABI/testing/sysfs-bus-pci#L344-L355
I0000 00:00:1723685861.119808   44933 cuda_executor.cc:1015] successful NUMA node read from SysFS had negative value (-1), but there must be at least one NUMA node, so returning NUMA node zero. See more at https://github.com/torvalds/linux/blob/v6.0/Documentation/ABI/testing/sysfs-bus-pci#L344-L355
I0000 00:00:1723685861.158075   44933 cuda_executor.cc:1015] successful NUMA node read from SysFS had negative value (-1), but there must be at least one NUMA node, so returning NUMA node zero. See more at https://github.com/torvalds/linux/blob/v6.0/Documentation/ABI/testing/sysfs-bus-pci#L344-L355
I0000 00:00:1723685861.160123   44933 cuda_executor.cc:1015] successful NUMA node read from SysFS had negative value (-1), but there must be at least one NUMA node, so returning NUMA node zero. See more at https://github.com/torvalds/linux/blob/v6.0/Documentation/ABI/testing/sysfs-bus-pci#L344-L355
I0000 00:00:1723685861.162060   44933 cuda_executor.cc:1015] successful NUMA node read from SysFS had negative value (-1), but there must be at least one NUMA node, so returning NUMA node zero. See more at https://github.com/torvalds/linux/blob/v6.0/Documentation/ABI/testing/sysfs-bus-pci#L344-L355
I0000 00:00:1723685861.163993   44933 cuda_executor.cc:1015] successful NUMA node read from SysFS had negative value (-1), but there must be at least one NUMA node, so returning NUMA node zero. See more at https://github.com/torvalds/linux/blob/v6.0/Documentation/ABI/testing/sysfs-bus-pci#L344-L355
I0000 00:00:1723685861.165963   44933 cuda_executor.cc:1015] successful NUMA node read from SysFS had negative value (-1), but there must be at least one NUMA node, so returning NUMA node zero. See more at https://github.com/torvalds/linux/blob/v6.0/Documentation/ABI/testing/sysfs-bus-pci#L344-L355
I0000 00:00:1723685861.167940   44933 cuda_executor.cc:1015] successful NUMA node read from SysFS had negative value (-1), but there must be at least one NUMA node, so returning NUMA node zero. See more at https://github.com/torvalds/linux/blob/v6.0/Documentation/ABI/testing/sysfs-bus-pci#L344-L355
I0000 00:00:1723685861.169863   44933 cuda_executor.cc:1015] successful NUMA node read from SysFS had negative value (-1), but there must be at least one NUMA node, so returning NUMA node zero. See more at https://github.com/torvalds/linux/blob/v6.0/Documentation/ABI/testing/sysfs-bus-pci#L344-L355
I0000 00:00:1723685861.171778   44933 cuda_executor.cc:1015] successful NUMA node read from SysFS had negative value (-1), but there must be at least one NUMA node, so returning NUMA node zero. See more at https://github.com/torvalds/linux/blob/v6.0/Documentation/ABI/testing/sysfs-bus-pci#L344-L355
I0000 00:00:1723685861.173638   44933 cuda_executor.cc:1015] successful NUMA node read from SysFS had negative value (-1), but there must be at least one NUMA node, so returning NUMA node zero. See more at https://github.com/torvalds/linux/blob/v6.0/Documentation/ABI/testing/sysfs-bus-pci#L344-L355
I0000 00:00:1723685861.176135   44933 cuda_executor.cc:1015] successful NUMA node read from SysFS had negative value (-1), but there must be at least one NUMA node, so returning NUMA node zero. See more at https://github.com/torvalds/linux/blob/v6.0/Documentation/ABI/testing/sysfs-bus-pci#L344-L355
I0000 00:00:1723685861.178420   44933 cuda_executor.cc:1015] successful NUMA node read from SysFS had negative value (-1), but there must be at least one NUMA node, so returning NUMA node zero. See more at https://github.com/torvalds/linux/blob/v6.0/Documentation/ABI/testing/sysfs-bus-pci#L344-L355
I0000 00:00:1723685861.180782   44933 cuda_executor.cc:1015] successful NUMA node read from SysFS had negative value (-1), but there must be at least one NUMA node, so returning NUMA node zero. See more at https://github.com/torvalds/linux/blob/v6.0/Documentation/ABI/testing/sysfs-bus-pci#L344-L355
<_TensorSliceDataset element_spec=TensorSpec(shape=(), dtype=tf.int32, name=None)>
for elem in dataset:
  print(elem.numpy())
8
3
0
8
2
1

或者透過使用 iter 顯式建立一個 Python 迭代器並使用 next 消耗其元素

it = iter(dataset)

print(next(it).numpy())
8

另外,資料集元素可以使用 reduce 轉換來消耗,該轉換將所有元素歸約以產生單個結果。以下示例說明了如何使用 reduce 轉換來計算整數資料集的和。

print(dataset.reduce(0, lambda state, value: state + value).numpy())
22

資料集結構

資料集產生一系列元素,其中每個元素都是相同的(巢狀)分量結構。結構中的各個分量可以是 tf.TypeSpec 可表示的任何型別,包括 tf.Tensortf.sparse.SparseTensortf.RaggedTensortf.TensorArraytf.data.Dataset

可用於表示元素(巢狀)結構的 Python 構造包括 tupledictNamedTupleOrderedDict。特別地,list 不是表示資料集元素結構的有效構造。這是因為早期的 tf.data 使用者強烈要求 list 輸入(例如,傳遞給 tf.data.Dataset.from_tensors 時)被自動打包為張量,而 list 輸出(例如,使用者定義函式的返回值)被強制轉換為 tuple。因此,如果您希望將 list 輸入視為一種結構,則需要將其轉換為 tuple;如果您希望 list 輸出成為單個分量,則需要使用 tf.stack 顯式對其進行打包。

Dataset.element_spec 屬性允許您檢查每個元素分量的型別。該屬性返回一個 tf.TypeSpec 物件的巢狀結構,該結構與元素的結構相匹配,元素可以是單個分量、分量元組或分量巢狀元組。例如

dataset1 = tf.data.Dataset.from_tensor_slices(tf.random.uniform([4, 10]))

dataset1.element_spec
TensorSpec(shape=(10,), dtype=tf.float32, name=None)
dataset2 = tf.data.Dataset.from_tensor_slices(
   (tf.random.uniform([4]),
    tf.random.uniform([4, 100], maxval=100, dtype=tf.int32)))

dataset2.element_spec
(TensorSpec(shape=(), dtype=tf.float32, name=None),
 TensorSpec(shape=(100,), dtype=tf.int32, name=None))
dataset3 = tf.data.Dataset.zip((dataset1, dataset2))

dataset3.element_spec
(TensorSpec(shape=(10,), dtype=tf.float32, name=None),
 (TensorSpec(shape=(), dtype=tf.float32, name=None),
  TensorSpec(shape=(100,), dtype=tf.int32, name=None)))
# Dataset containing a sparse tensor.
dataset4 = tf.data.Dataset.from_tensors(tf.SparseTensor(indices=[[0, 0], [1, 2]], values=[1, 2], dense_shape=[3, 4]))

dataset4.element_spec
SparseTensorSpec(TensorShape([3, 4]), tf.int32)
# Use value_type to see the type of value represented by the element spec
dataset4.element_spec.value_type
tensorflow.python.framework.sparse_tensor.SparseTensor

Dataset 轉換支援任何結構的資料集。當使用對每個元素應用函式的 Dataset.mapDataset.filter 轉換時,元素結構決定了函式的引數

dataset1 = tf.data.Dataset.from_tensor_slices(
    tf.random.uniform([4, 10], minval=1, maxval=10, dtype=tf.int32))

dataset1
<_TensorSliceDataset element_spec=TensorSpec(shape=(10,), dtype=tf.int32, name=None)>
for z in dataset1:
  print(z.numpy())
[3 4 1 6 1 8 5 8 9 4]
[2 7 6 9 2 6 6 4 9 7]
[8 7 9 6 3 4 5 8 4 4]
[2 1 1 1 3 9 7 8 6 8]
dataset2 = tf.data.Dataset.from_tensor_slices(
   (tf.random.uniform([4]),
    tf.random.uniform([4, 100], maxval=100, dtype=tf.int32)))

dataset2
<_TensorSliceDataset element_spec=(TensorSpec(shape=(), dtype=tf.float32, name=None), TensorSpec(shape=(100,), dtype=tf.int32, name=None))>
dataset3 = tf.data.Dataset.zip((dataset1, dataset2))

dataset3
<_ZipDataset element_spec=(TensorSpec(shape=(10,), dtype=tf.int32, name=None), (TensorSpec(shape=(), dtype=tf.float32, name=None), TensorSpec(shape=(100,), dtype=tf.int32, name=None)))>
for a, (b,c) in dataset3:
  print('shapes: {a.shape}, {b.shape}, {c.shape}'.format(a=a, b=b, c=c))
shapes: (10,), (), (100,)
shapes: (10,), (), (100,)
shapes: (10,), (), (100,)
shapes: (10,), (), (100,)

讀取輸入資料

消耗 NumPy 陣列

有關更多示例,請參閱載入 NumPy 陣列教程。

如果您的所有輸入資料都適合放入記憶體,則從中建立 Dataset 的最簡單方法是將它們轉換為 tf.Tensor 物件並使用 Dataset.from_tensor_slices

train, test = tf.keras.datasets.fashion_mnist.load_data()
Downloading data from https://storage.googleapis.com/tensorflow/tf-keras-datasets/train-labels-idx1-ubyte.gz
29515/29515 ━━━━━━━━━━━━━━━━━━━━ 0s 0us/step
Downloading data from https://storage.googleapis.com/tensorflow/tf-keras-datasets/train-images-idx3-ubyte.gz
26421880/26421880 ━━━━━━━━━━━━━━━━━━━━ 0s 0us/step
Downloading data from https://storage.googleapis.com/tensorflow/tf-keras-datasets/t10k-labels-idx1-ubyte.gz
5148/5148 ━━━━━━━━━━━━━━━━━━━━ 0s 0us/step
Downloading data from https://storage.googleapis.com/tensorflow/tf-keras-datasets/t10k-images-idx3-ubyte.gz
4422102/4422102 ━━━━━━━━━━━━━━━━━━━━ 0s 0us/step
images, labels = train
images = images/255

dataset = tf.data.Dataset.from_tensor_slices((images, labels))
dataset
<_TensorSliceDataset element_spec=(TensorSpec(shape=(28, 28), dtype=tf.float64, name=None), TensorSpec(shape=(), dtype=tf.uint8, name=None))>

消耗 Python 生成器

另一個可以輕鬆作為 tf.data.Dataset 獲取的常見資料來源是 Python 生成器。

def count(stop):
  i = 0
  while i<stop:
    yield i
    i += 1
for n in count(5):
  print(n)
0
1
2
3
4

Dataset.from_generator 建構函式將 Python 生成器轉換為功能完備的 tf.data.Dataset

建構函式接收一個可呼叫物件作為輸入,而不是迭代器。這允許它在生成器到達末尾時重新啟動生成器。它接受一個可選的 args 引數,該引數作為可呼叫物件的引數傳遞。

output_types 引數是必需的,因為 tf.data 在內部構建一個 tf.Graph,而圖邊需要一個 tf.dtype

ds_counter = tf.data.Dataset.from_generator(count, args=[25], output_types=tf.int32, output_shapes = (), )
for count_batch in ds_counter.repeat().batch(10).take(10):
  print(count_batch.numpy())
[0 1 2 3 4 5 6 7 8 9]
[10 11 12 13 14 15 16 17 18 19]
[20 21 22 23 24  0  1  2  3  4]
[ 5  6  7  8  9 10 11 12 13 14]
[15 16 17 18 19 20 21 22 23 24]
[0 1 2 3 4 5 6 7 8 9]
[10 11 12 13 14 15 16 17 18 19]
[20 21 22 23 24  0  1  2  3  4]
[ 5  6  7  8  9 10 11 12 13 14]
[15 16 17 18 19 20 21 22 23 24]

output_shapes 引數並非必需,但強烈建議使用,因為許多 TensorFlow 操作不支援秩未知的張量。如果特定軸的長度未知或可變,請在 output_shapes 中將其設定為 None

同樣需要注意的是,output_shapesoutput_types 遵循與其他資料集方法相同的巢狀規則。

這是一個演示這兩個方面的生成器示例:它返回陣列元組,其中第二個陣列是長度未知的向量。

def gen_series():
  i = 0
  while True:
    size = np.random.randint(0, 10)
    yield i, np.random.normal(size=(size,))
    i += 1
for i, series in gen_series():
  print(i, ":", str(series))
  if i > 5:
    break
0 : [1.1274]
1 : [-0.5822  0.8497 -1.3594  0.2083 -0.3007  1.2171 -0.3551]
2 : [-1.2016 -0.1085  0.4088  0.0801  1.4901 -2.3102]
3 : [ 0.5816 -0.6447 -0.9673  0.5282  0.52   -0.2634  0.3001  0.8753]
4 : [ 0.0888  0.071   1.26   -0.347  -0.2643 -1.0757  0.4192]
5 : [ 0.4911  0.8377  0.3576 -0.0351  0.9663]
6 : [-0.1996  0.5808  0.4589  1.8229 -0.5712]

第一個輸出是 int32,第二個是 float32

第一項是標量,形狀為 (),第二項是長度未知的向量,形狀為 (None,)

ds_series = tf.data.Dataset.from_generator(
    gen_series,
    output_types=(tf.int32, tf.float32),
    output_shapes=((), (None,)))

ds_series
<_FlatMapDataset element_spec=(TensorSpec(shape=(), dtype=tf.int32, name=None), TensorSpec(shape=(None,), dtype=tf.float32, name=None))>

現在它可以像普通的 tf.data.Dataset 一樣使用。請注意,在批處理具有可變形狀的資料集時,您需要使用 Dataset.padded_batch

ds_series_batch = ds_series.shuffle(20).padded_batch(10)

ids, sequence_batch = next(iter(ds_series_batch))
print(ids.numpy())
print()
print(sequence_batch.numpy())
[ 5 19 20 11  4 10 17  8 27 18]

[[-0.7479  0.867  -0.0558 -1.0825 -0.4113  0.0312  0.    ]
 [-1.0498 -0.3941  0.      0.      0.      0.      0.    ]
 [-0.2709  0.0236  0.0746  0.3704  0.      0.      0.    ]
 [ 1.6525 -0.861   0.5642  0.9961  0.7463  0.      0.    ]
 [ 0.4122 -0.118   1.5491  1.9578  0.      0.      0.    ]
 [-1.6237  1.3636 -0.2079  0.      0.      0.      0.    ]
 [ 0.      0.      0.      0.      0.      0.      0.    ]
 [ 0.      0.      0.      0.      0.      0.      0.    ]
 [-1.3268  0.9881  0.531   0.      0.      0.      0.    ]
 [ 0.0284 -1.4974 -0.545  -1.2795  0.7032  1.4058  0.1412]]

對於更現實的示例,請嘗試將 preprocessing.image.ImageDataGenerator 包裝為 tf.data.Dataset

首先下載資料

flowers = tf.keras.utils.get_file(
    'flower_photos',
    'https://storage.googleapis.com/download.tensorflow.org/example_images/flower_photos.tgz',
    untar=True)
Downloading data from https://storage.googleapis.com/download.tensorflow.org/example_images/flower_photos.tgz
228813984/228813984 ━━━━━━━━━━━━━━━━━━━━ 2s 0us/step

建立 image.ImageDataGenerator

img_gen = tf.keras.preprocessing.image.ImageDataGenerator(rescale=1./255, rotation_range=20)
images, labels = next(img_gen.flow_from_directory(flowers))
Found 3670 images belonging to 5 classes.
print(images.dtype, images.shape)
print(labels.dtype, labels.shape)
float32 (32, 256, 256, 3)
float32 (32, 5)
ds = tf.data.Dataset.from_generator(
    lambda: img_gen.flow_from_directory(flowers),
    output_types=(tf.float32, tf.float32),
    output_shapes=([32,256,256,3], [32,5])
)

ds.element_spec
(TensorSpec(shape=(32, 256, 256, 3), dtype=tf.float32, name=None),
 TensorSpec(shape=(32, 5), dtype=tf.float32, name=None))
for images, labels in ds.take(1):
  print('images.shape: ', images.shape)
  print('labels.shape: ', labels.shape)
Found 3670 images belonging to 5 classes.
images.shape:  (32, 256, 256, 3)
labels.shape:  (32, 5)

消耗 TFRecord 資料

有關端到端示例,請參閱載入 TFRecords 教程。

tf.data API 支援多種檔案格式,以便您可以處理無法放入記憶體的大型資料集。例如,TFRecord 檔案格式是一種簡單的面向記錄的二進位制格式,許多 TensorFlow 應用程式將其用於訓練資料。tf.data.TFRecordDataset 類使您能夠作為輸入流水線的一部分流式處理一個或多個 TFRecord 檔案的內容。

這是一個使用來自法國街道名稱標誌 (FSNS) 的測試檔案的示例。

# Creates a dataset that reads all of the examples from two files.
fsns_test_file = tf.keras.utils.get_file("fsns.tfrec", "https://storage.googleapis.com/download.tensorflow.org/data/fsns-20160927/testdata/fsns-00000-of-00001")
Downloading data from https://storage.googleapis.com/download.tensorflow.org/data/fsns-20160927/testdata/fsns-00000-of-00001
7904079/7904079 ━━━━━━━━━━━━━━━━━━━━ 0s 0us/step

TFRecordDataset 初始化程式的 filenames 引數可以是字串、字串列表或字串 tf.Tensor。因此,如果您有兩組用於訓練和驗證目的的檔案,則可以建立一個工廠方法來生成資料集,並將檔名作為輸入引數

dataset = tf.data.TFRecordDataset(filenames = [fsns_test_file])
dataset
<TFRecordDatasetV2 element_spec=TensorSpec(shape=(), dtype=tf.string, name=None)>

許多 TensorFlow 專案在 TFRecord 檔案中使用序列化的 tf.train.Example 記錄。這些在檢查之前需要解碼

raw_example = next(iter(dataset))
parsed = tf.train.Example.FromString(raw_example.numpy())

parsed.features.feature['image/text']
bytes_list {
  value: "Rue Perreyon"
}

消耗文字資料

有關端到端示例,請參閱載入文字教程。

許多資料集作為一種或多種文字檔案分發。tf.data.TextLineDataset 提供了一種從一個或多個文字檔案中提取行的簡便方法。給定一個或多個檔名,TextLineDataset 將為這些檔案的每一行生成一個字串值元素。

directory_url = 'https://storage.googleapis.com/download.tensorflow.org/data/illiad/'
file_names = ['cowper.txt', 'derby.txt', 'butler.txt']

file_paths = [
    tf.keras.utils.get_file(file_name, directory_url + file_name)
    for file_name in file_names
]
Downloading data from https://storage.googleapis.com/download.tensorflow.org/data/illiad/cowper.txt
815980/815980 ━━━━━━━━━━━━━━━━━━━━ 0s 0us/step
Downloading data from https://storage.googleapis.com/download.tensorflow.org/data/illiad/derby.txt
809730/809730 ━━━━━━━━━━━━━━━━━━━━ 0s 0us/step
Downloading data from https://storage.googleapis.com/download.tensorflow.org/data/illiad/butler.txt
807992/807992 ━━━━━━━━━━━━━━━━━━━━ 0s 0us/step
dataset = tf.data.TextLineDataset(file_paths)

這是第一個檔案的前幾行

for line in dataset.take(5):
  print(line.numpy())
b"\xef\xbb\xbfAchilles sing, O Goddess! Peleus' son;"
b'His wrath pernicious, who ten thousand woes'
b"Caused to Achaia's host, sent many a soul"
b'Illustrious into Ades premature,'
b'And Heroes gave (so stood the will of Jove)'

要交替顯示檔案之間的行,請使用 Dataset.interleave。這使得更容易將檔案混合在一起。這是每次翻譯的第一行、第二行和第三行

files_ds = tf.data.Dataset.from_tensor_slices(file_paths)
lines_ds = files_ds.interleave(tf.data.TextLineDataset, cycle_length=3)

for i, line in enumerate(lines_ds.take(9)):
  if i % 3 == 0:
    print()
  print(line.numpy())
b"\xef\xbb\xbfAchilles sing, O Goddess! Peleus' son;"
b"\xef\xbb\xbfOf Peleus' son, Achilles, sing, O Muse,"
b'\xef\xbb\xbfSing, O goddess, the anger of Achilles son of Peleus, that brought'

b'His wrath pernicious, who ten thousand woes'
b'The vengeance, deep and deadly; whence to Greece'
b'countless ills upon the Achaeans. Many a brave soul did it send'

b"Caused to Achaia's host, sent many a soul"
b'Unnumbered ills arose; which many a soul'
b'hurrying down to Hades, and many a hero did it yield a prey to dogs and'

預設情況下,TextLineDataset 會產生每個檔案的每一行,這可能不是我們想要的,例如,如果檔案以標題行開頭或包含註釋。可以使用 Dataset.skip()Dataset.filter 轉換來刪除這些行。在這裡,您跳過第一行,然後進行過濾以僅找到倖存者。

titanic_file = tf.keras.utils.get_file("train.csv", "https://storage.googleapis.com/tf-datasets/titanic/train.csv")
titanic_lines = tf.data.TextLineDataset(titanic_file)
Downloading data from https://storage.googleapis.com/tf-datasets/titanic/train.csv
30874/30874 ━━━━━━━━━━━━━━━━━━━━ 0s 0us/step
for line in titanic_lines.take(10):
  print(line.numpy())
b'survived,sex,age,n_siblings_spouses,parch,fare,class,deck,embark_town,alone'
b'0,male,22.0,1,0,7.25,Third,unknown,Southampton,n'
b'1,female,38.0,1,0,71.2833,First,C,Cherbourg,n'
b'1,female,26.0,0,0,7.925,Third,unknown,Southampton,y'
b'1,female,35.0,1,0,53.1,First,C,Southampton,n'
b'0,male,28.0,0,0,8.4583,Third,unknown,Queenstown,y'
b'0,male,2.0,3,1,21.075,Third,unknown,Southampton,n'
b'1,female,27.0,0,2,11.1333,Third,unknown,Southampton,n'
b'1,female,14.0,1,0,30.0708,Second,unknown,Cherbourg,n'
b'1,female,4.0,1,1,16.7,Third,G,Southampton,n'
def survived(line):
  return tf.not_equal(tf.strings.substr(line, 0, 1), "0")

survivors = titanic_lines.skip(1).filter(survived)
for line in survivors.take(10):
  print(line.numpy())
b'1,female,38.0,1,0,71.2833,First,C,Cherbourg,n'
b'1,female,26.0,0,0,7.925,Third,unknown,Southampton,y'
b'1,female,35.0,1,0,53.1,First,C,Southampton,n'
b'1,female,27.0,0,2,11.1333,Third,unknown,Southampton,n'
b'1,female,14.0,1,0,30.0708,Second,unknown,Cherbourg,n'
b'1,female,4.0,1,1,16.7,Third,G,Southampton,n'
b'1,male,28.0,0,0,13.0,Second,unknown,Southampton,y'
b'1,female,28.0,0,0,7.225,Third,unknown,Cherbourg,y'
b'1,male,28.0,0,0,35.5,First,A,Southampton,y'
b'1,female,38.0,1,5,31.3875,Third,unknown,Southampton,n'

消耗 CSV 資料

有關更多示例,請參閱載入 CSV 檔案載入 Pandas DataFrames 教程。

CSV 檔案格式是一種用於以純文字儲存表格資料的流行格式。

例如:

titanic_file = tf.keras.utils.get_file("train.csv", "https://storage.googleapis.com/tf-datasets/titanic/train.csv")
df = pd.read_csv(titanic_file)
df.head()

如果您的資料適合放入記憶體,則同樣的 Dataset.from_tensor_slices 方法適用於字典,從而使此資料易於匯入

titanic_slices = tf.data.Dataset.from_tensor_slices(dict(df))

for feature_batch in titanic_slices.take(1):
  for key, value in feature_batch.items():
    print("  {!r:20s}: {}".format(key, value))
'survived'          : 0
  'sex'               : b'male'
  'age'               : 22.0
  'n_siblings_spouses': 1
  'parch'             : 0
  'fare'              : 7.25
  'class'             : b'Third'
  'deck'              : b'unknown'
  'embark_town'       : b'Southampton'
  'alone'             : b'n'

一種更具可擴充套件性的方法是根據需要從磁碟載入。

tf.data 模組提供了從符合 RFC 4180 的一個或多個 CSV 檔案中提取記錄的方法。

tf.data.experimental.make_csv_dataset 函式是用於讀取一組 CSV 檔案的高階介面。它支援列型別推斷和許多其他功能(如批處理和重排),以簡化使用。

titanic_batches = tf.data.experimental.make_csv_dataset(
    titanic_file, batch_size=4,
    label_name="survived")
for feature_batch, label_batch in titanic_batches.take(1):
  print("'survived': {}".format(label_batch))
  print("features:")
  for key, value in feature_batch.items():
    print("  {!r:20s}: {}".format(key, value))
'survived': [0 0 0 0]
features:
  'sex'               : [b'male' b'male' b'male' b'male']
  'age'               : [28. 46. 28. 26.]
  'n_siblings_spouses': [0 1 0 0]
  'parch'             : [1 0 0 0]
  'fare'              : [33.     61.175   8.05    7.8875]
  'class'             : [b'Second' b'First' b'Third' b'Third']
  'deck'              : [b'unknown' b'E' b'unknown' b'unknown']
  'embark_town'       : [b'Southampton' b'Southampton' b'Southampton' b'Southampton']
  'alone'             : [b'n' b'n' b'y' b'y']

如果您只需要列的子集,可以使用 select_columns 引數。

titanic_batches = tf.data.experimental.make_csv_dataset(
    titanic_file, batch_size=4,
    label_name="survived", select_columns=['class', 'fare', 'survived'])
for feature_batch, label_batch in titanic_batches.take(1):
  print("'survived': {}".format(label_batch))
  for key, value in feature_batch.items():
    print("  {!r:20s}: {}".format(key, value))
'survived': [1 1 0 0]
  'fare'              : [10.5   35.5   12.875 29.125]
  'class'             : [b'Second' b'First' b'Second' b'Third']

還有一個較低級別的 experimental.CsvDataset 類,它提供更細粒度的控制。它不支援列型別推斷。相反,您必須指定每一列的型別。

titanic_types  = [tf.int32, tf.string, tf.float32, tf.int32, tf.int32, tf.float32, tf.string, tf.string, tf.string, tf.string]
dataset = tf.data.experimental.CsvDataset(titanic_file, titanic_types , header=True)

for line in dataset.take(10):
  print([item.numpy() for item in line])
[0, b'male', 22.0, 1, 0, 7.25, b'Third', b'unknown', b'Southampton', b'n']
[1, b'female', 38.0, 1, 0, 71.2833, b'First', b'C', b'Cherbourg', b'n']
[1, b'female', 26.0, 0, 0, 7.925, b'Third', b'unknown', b'Southampton', b'y']
[1, b'female', 35.0, 1, 0, 53.1, b'First', b'C', b'Southampton', b'n']
[0, b'male', 28.0, 0, 0, 8.4583, b'Third', b'unknown', b'Queenstown', b'y']
[0, b'male', 2.0, 3, 1, 21.075, b'Third', b'unknown', b'Southampton', b'n']
[1, b'female', 27.0, 0, 2, 11.1333, b'Third', b'unknown', b'Southampton', b'n']
[1, b'female', 14.0, 1, 0, 30.0708, b'Second', b'unknown', b'Cherbourg', b'n']
[1, b'female', 4.0, 1, 1, 16.7, b'Third', b'G', b'Southampton', b'n']
[0, b'male', 20.0, 0, 0, 8.05, b'Third', b'unknown', b'Southampton', b'y']

如果某些列為空,此低階介面允許您提供預設值而不是列型別。

%%writefile missing.csv
1,2,3,4
,2,3,4
1,,3,4
1,2,,4
1,2,3,
,,,
Writing missing.csv
# Creates a dataset that reads all of the records from two CSV files, each with
# four float columns which may have missing values.

record_defaults = [999,999,999,999]
dataset = tf.data.experimental.CsvDataset("missing.csv", record_defaults)
dataset = dataset.map(lambda *items: tf.stack(items))
dataset
<_MapDataset element_spec=TensorSpec(shape=(4,), dtype=tf.int32, name=None)>
for line in dataset:
  print(line.numpy())
[1 2 3 4]
[999   2   3   4]
[  1 999   3   4]
[  1   2 999   4]
[  1   2   3 999]
[999 999 999 999]

預設情況下,CsvDataset 會產生檔案中每一行的每一列,這可能不是我們想要的,例如如果檔案以應該忽略的標題行開頭,或者如果某些列在輸入中不需要。這些行和欄位可以分別透過 headerselect_cols 引數刪除。

# Creates a dataset that reads all of the records from two CSV files with
# headers, extracting float data from columns 2 and 4.
record_defaults = [999, 999] # Only provide defaults for the selected columns
dataset = tf.data.experimental.CsvDataset("missing.csv", record_defaults, select_cols=[1, 3])
dataset = dataset.map(lambda *items: tf.stack(items))
dataset
<_MapDataset element_spec=TensorSpec(shape=(2,), dtype=tf.int32, name=None)>
for line in dataset:
  print(line.numpy())
[2 4]
[2 4]
[999   4]
[2 4]
[  2 999]
[999 999]

消耗一組檔案

有許多資料集作為一組檔案分發,其中每個檔案都是一個示例。

flowers_root = tf.keras.utils.get_file(
    'flower_photos',
    'https://storage.googleapis.com/download.tensorflow.org/example_images/flower_photos.tgz',
    untar=True)
flowers_root = pathlib.Path(flowers_root)

根目錄包含每個類別的目錄

for item in flowers_root.glob("*"):
  print(item.name)
daisy
tulips
sunflowers
LICENSE.txt
dandelion
roses

每個類別目錄中的檔案都是示例

list_ds = tf.data.Dataset.list_files(str(flowers_root/'*/*'))

for f in list_ds.take(5):
  print(f.numpy())
b'/home/kbuilder/.keras/datasets/flower_photos/tulips/4955884820_7e4ce4d7e5_m.jpg'
b'/home/kbuilder/.keras/datasets/flower_photos/dandelion/6250363717_17732e992e_n.jpg'
b'/home/kbuilder/.keras/datasets/flower_photos/tulips/14278331403_4c475f9a9b.jpg'
b'/home/kbuilder/.keras/datasets/flower_photos/dandelion/480621885_4c8b50fa11_m.jpg'
b'/home/kbuilder/.keras/datasets/flower_photos/tulips/5716293002_a8be6a6dd3_n.jpg'

使用 tf.io.read_file 函式讀取資料並從路徑中提取標籤,返回 (image, label)

def process_path(file_path):
  label = tf.strings.split(file_path, os.sep)[-2]
  return tf.io.read_file(file_path), label

labeled_ds = list_ds.map(process_path)
for image_raw, label_text in labeled_ds.take(1):
  print(repr(image_raw.numpy()[:100]))
  print()
  print(label_text.numpy())
b'\xff\xd8\xff\xe0\x00\x10JFIF\x00\x01\x01\x00\x00\x01\x00\x01\x00\x00\xff\xdb\x00C\x00\x03\x02\x02\x03\x02\x02\x03\x03\x03\x03\x04\x03\x03\x04\x05\x08\x05\x05\x04\x04\x05\n\x07\x07\x06\x08\x0c\n\x0c\x0c\x0b\n\x0b\x0b\r\x0e\x12\x10\r\x0e\x11\x0e\x0b\x0b\x10\x16\x10\x11\x13\x14\x15\x15\x15\x0c\x0f\x17\x18\x16\x14\x18\x12\x14\x15\x14\xff\xdb\x00C\x01\x03\x04\x04\x05\x04\x05'

b'dandelion'

批處理資料集元素

簡單批處理

批處理的最簡單形式是將資料集的 n 個連續元素堆疊為一個元素。Dataset.batch() 轉換正是這樣做的,它具有與 tf.stack() 運算元相同的約束,並應用於元素的每個分量:即,對於每個分量 i,所有元素都必須具有完全相同形狀的張量。

inc_dataset = tf.data.Dataset.range(100)
dec_dataset = tf.data.Dataset.range(0, -100, -1)
dataset = tf.data.Dataset.zip((inc_dataset, dec_dataset))
batched_dataset = dataset.batch(4)

for batch in batched_dataset.take(4):
  print([arr.numpy() for arr in batch])
[array([0, 1, 2, 3]), array([ 0, -1, -2, -3])]
[array([4, 5, 6, 7]), array([-4, -5, -6, -7])]
[array([ 8,  9, 10, 11]), array([ -8,  -9, -10, -11])]
[array([12, 13, 14, 15]), array([-12, -13, -14, -15])]

雖然 tf.data 會嘗試傳播形狀資訊,但 Dataset.batch 的預設設定會導致批處理大小未知,因為最後一個批次可能不是滿的。注意形狀中的 None

batched_dataset
<_BatchDataset element_spec=(TensorSpec(shape=(None,), dtype=tf.int64, name=None), TensorSpec(shape=(None,), dtype=tf.int64, name=None))>

使用 drop_remainder 引數來忽略最後一個批次,並獲得完全的形狀傳播

batched_dataset = dataset.batch(7, drop_remainder=True)
batched_dataset
<_BatchDataset element_spec=(TensorSpec(shape=(7,), dtype=tf.int64, name=None), TensorSpec(shape=(7,), dtype=tf.int64, name=None))>

帶填充的張量批處理

上述方案適用於所有大小相同的張量。然而,許多模型(包括序列模型)使用大小可變的輸入資料(例如不同長度的序列)。為了處理這種情況,Dataset.padded_batch 轉換使您能夠透過指定可以填充的一個或多個維度來對不同形狀的張量進行批處理。

dataset = tf.data.Dataset.range(100)
dataset = dataset.map(lambda x: tf.fill([tf.cast(x, tf.int32)], x))
dataset = dataset.padded_batch(4, padded_shapes=(None,))

for batch in dataset.take(2):
  print(batch.numpy())
  print()
[[0 0 0]
 [1 0 0]
 [2 2 0]
 [3 3 3]]

[[4 4 4 4 0 0 0]
 [5 5 5 5 5 0 0]
 [6 6 6 6 6 6 0]
 [7 7 7 7 7 7 7]]

Dataset.padded_batch 轉換允許您為每個分量的每個維度設定不同的填充,並且它可以是可變長度的(由上面的示例中的 None 表示)或恆定長度。也可以覆蓋預設為 0 的填充值。

訓練工作流

處理多個 epoch

tf.data API 提供了兩種主要方法來處理相同資料的多個 epoch。

在多個 epoch 中迭代資料集的最簡單方法是使用 Dataset.repeat() 轉換。首先,建立一個泰坦尼克號資料集

titanic_file = tf.keras.utils.get_file("train.csv", "https://storage.googleapis.com/tf-datasets/titanic/train.csv")
titanic_lines = tf.data.TextLineDataset(titanic_file)
def plot_batch_sizes(ds):
  batch_sizes = [batch.shape[0] for batch in ds]
  plt.bar(range(len(batch_sizes)), batch_sizes)
  plt.xlabel('Batch number')
  plt.ylabel('Batch size')

應用不帶引數的 Dataset.repeat() 轉換將無限期地重複輸入。

Dataset.repeat 轉換會拼接其引數,而不會發出一個 epoch 結束和下一個 epoch 開始的訊號。因此,在 Dataset.repeat 之後應用的 Dataset.batch 將產生跨越 epoch 邊界的批次

titanic_batches = titanic_lines.repeat(3).batch(128)
plot_batch_sizes(titanic_batches)

png

如果您需要明確的 epoch 分隔,請將 Dataset.batch 放在 repeat 之前

titanic_batches = titanic_lines.batch(128).repeat(3)

plot_batch_sizes(titanic_batches)

png

如果您想在每個 epoch 結束時執行自定義計算(例如,收集統計資訊),那麼最簡單的方法是在每個 epoch 重新開始資料集迭代

epochs = 3
dataset = titanic_lines.batch(128)

for epoch in range(epochs):
  for batch in dataset:
    print(batch.shape)
  print("End of epoch: ", epoch)
(128,)
(128,)
(128,)
(128,)
(116,)
End of epoch:  0
(128,)
(128,)
(128,)
(128,)
(116,)
End of epoch:  1
(128,)
(128,)
(128,)
(128,)
(116,)
End of epoch:  2

隨機重排輸入資料

Dataset.shuffle() 轉換維護一個固定大小的緩衝區,並從該緩衝區中均勻隨機地選擇下一個元素。

向資料集新增索引,以便您可以看到效果

lines = tf.data.TextLineDataset(titanic_file)
counter = tf.data.experimental.Counter()

dataset = tf.data.Dataset.zip((counter, lines))
dataset = dataset.shuffle(buffer_size=100)
dataset = dataset.batch(20)
dataset
WARNING:tensorflow:From /tmpfs/tmp/ipykernel_44933/4092668703.py:2: CounterV2 (from tensorflow.python.data.experimental.ops.counter) is deprecated and will be removed in a future version.
Instructions for updating:
Use `tf.data.Dataset.counter(...)` instead.
<_BatchDataset element_spec=(TensorSpec(shape=(None,), dtype=tf.int64, name=None), TensorSpec(shape=(None,), dtype=tf.string, name=None))>

由於 buffer_size 為 100,批大小為 20,因此第一個批次不包含索引超過 120 的元素。

n,line_batch = next(iter(dataset))
print(n.numpy())
[ 99  18   1  29  66  88  47  30  80  46  68  44  35  40  33  95 108 105
  38 113]

Dataset.batch 一樣,相對於 Dataset.repeat 的順序很重要。

Dataset.shuffle 不會發出 epoch 結束的訊號,直到重排緩衝區為空。因此,放置在 repeat 之前的 shuffle 將顯示一個 epoch 的所有元素,然後移動到下一個

dataset = tf.data.Dataset.zip((counter, lines))
shuffled = dataset.shuffle(buffer_size=100).batch(10).repeat(2)

print("Here are the item ID's near the epoch boundary:\n")
for n, line_batch in shuffled.skip(60).take(5):
  print(n.numpy())
Here are the item ID's near the epoch boundary:

[469 414 497 584 615 612 625 627 603 621]
[582 553 343 602 626 567 486 593 616 525]
[557 576 478 533 591 398 484 431]
[66 43 51 18 94  3 76 52 90 57]
[  0 101  71  86  56  17  33  70 110  75]
shuffle_repeat = [n.numpy().mean() for n, line_batch in shuffled]
plt.plot(shuffle_repeat, label="shuffle().repeat()")
plt.ylabel("Mean item ID")
plt.legend()
<matplotlib.legend.Legend at 0x7f373c471af0>

png

但是 repeat 放在 shuffle 之前會將 epoch 邊界混合在一起

dataset = tf.data.Dataset.zip((counter, lines))
shuffled = dataset.repeat(2).shuffle(buffer_size=100).batch(10)

print("Here are the item ID's near the epoch boundary:\n")
for n, line_batch in shuffled.skip(55).take(15):
  print(n.numpy())
Here are the item ID's near the epoch boundary:

[583 415   1 542 563   9 620 622 551 548]
[589 592 365 571  33 557 618  31 541  27]
[537  24 615  43  18 550  11   8  39 369]
[601  38 485  20 627  46  22  23 322 608]
[626 590 491  63  29 564  17  19 617  66]
[508 580  72  45  57  54 556  62  14 511]
[623  73  75  79 599 372  21  83 547  26]
[486   4   0 573  74  49   0  53  95  34]
[ 60 605  15  90  99 549  16  50  91  80]
[106 108 112 297 561  44  52  82  86  71]
[581  77 117  28 567  10  30   3  81  89]
[587  32 102   7 135  51 113 110 114 451]
[ 59  64  68 116  76 306 367 128 552 136]
[111 569 522   5  67 616 154 131 512  37]
[539 103 142  78  85   2  87  12 149 137]
repeat_shuffle = [n.numpy().mean() for n, line_batch in shuffled]

plt.plot(shuffle_repeat, label="shuffle().repeat()")
plt.plot(repeat_shuffle, label="repeat().shuffle()")
plt.ylabel("Mean item ID")
plt.legend()
<matplotlib.legend.Legend at 0x7f373c462b20>

png

預處理資料

Dataset.map(f) 轉換透過對輸入資料集的每個元素應用給定的函式 f 來生成新的資料集。它基於函數語言程式設計語言中通常應用於列表(和其他結構)的 map() 函式。函式 f 接收表示輸入中單個元素的 tf.Tensor 物件,並返回將表示新資料集中單個元素的 tf.Tensor 物件。其實現使用標準 TensorFlow 操作將一個元素轉換為另一個元素。

本節涵蓋了如何使用 Dataset.map() 的常見示例。

解碼影像資料並調整其大小

在真實世界的影像資料上訓練神經網路時,通常有必要將不同大小的影像轉換為通用大小,以便它們可以被批處理為固定大小。

重建花卉檔名資料集

list_ds = tf.data.Dataset.list_files(str(flowers_root/'*/*'))

編寫一個操作資料集元素的函式。

# Reads an image from a file, decodes it into a dense tensor, and resizes it
# to a fixed shape.
def parse_image(filename):
  parts = tf.strings.split(filename, os.sep)
  label = parts[-2]

  image = tf.io.read_file(filename)
  image = tf.io.decode_jpeg(image)
  image = tf.image.convert_image_dtype(image, tf.float32)
  image = tf.image.resize(image, [128, 128])
  return image, label

測試它是否有效。

file_path = next(iter(list_ds))
image, label = parse_image(file_path)

def show(image, label):
  plt.figure()
  plt.imshow(image)
  plt.title(label.numpy().decode('utf-8'))
  plt.axis('off')

show(image, label)

png

將其對映到資料集上。

images_ds = list_ds.map(parse_image)

for image, label in images_ds.take(2):
  show(image, label)

png

png

應用任意 Python 邏輯

出於效能原因,請儘可能使用 TensorFlow 操作來預處理資料。但是,在解析輸入資料時呼叫外部 Python 庫有時很有用。您可以在 Dataset.map 轉換中使用 tf.py_function 操作。

例如,如果您想應用隨機旋轉,tf.image 模組只有 tf.image.rot90,這對影像增強用處不大。

為了演示 tf.py_function,請嘗試改用 scipy.ndimage.rotate 函式

import scipy.ndimage as ndimage

@tf.py_function(Tout=tf.float32)
def random_rotate_image(image):
  image = ndimage.rotate(image, np.random.uniform(-30, 30), reshape=False)
  return image
image, label = next(iter(images_ds))
image = random_rotate_image(image)
show(image, label)
Clipping input data to the valid range for imshow with RGB data ([0..1] for floats or [0..255] for integers). Got range [-0.07214577..1.0803627].

png

要將此函式與 Dataset.map 一起使用,適用的注意事項與 Dataset.from_generator 相同,您需要在應用函式時描述返回形狀和型別

def tf_random_rotate_image(image, label):
  im_shape = image.shape
  image = random_rotate_image(image)
  image.set_shape(im_shape)
  return image, label
rot_ds = images_ds.map(tf_random_rotate_image)

for image, label in rot_ds.take(2):
  show(image, label)
Clipping input data to the valid range for imshow with RGB data ([0..1] for floats or [0..255] for integers). Got range [-0.014158356..1.0156134].
Clipping input data to the valid range for imshow with RGB data ([0..1] for floats or [0..255] for integers). Got range [-0.067302234..1.1018459].

png

png

解析 tf.Example 協議緩衝區訊息

許多輸入流水線從 TFRecord 格式中提取 tf.train.Example 協議緩衝區訊息。每條 tf.train.Example 記錄都包含一個或多個“特徵”,輸入流水線通常會將這些特徵轉換為張量。

fsns_test_file = tf.keras.utils.get_file("fsns.tfrec", "https://storage.googleapis.com/download.tensorflow.org/data/fsns-20160927/testdata/fsns-00000-of-00001")
dataset = tf.data.TFRecordDataset(filenames = [fsns_test_file])
dataset
<TFRecordDatasetV2 element_spec=TensorSpec(shape=(), dtype=tf.string, name=None)>

您可以在 tf.data.Dataset 之外使用 tf.train.Example 原型來理解資料

raw_example = next(iter(dataset))
parsed = tf.train.Example.FromString(raw_example.numpy())

feature = parsed.features.feature
raw_img = feature['image/encoded'].bytes_list.value[0]
img = tf.image.decode_png(raw_img)
plt.imshow(img)
plt.axis('off')
_ = plt.title(feature["image/text"].bytes_list.value[0])

png

raw_example = next(iter(dataset))
def tf_parse(eg):
  example = tf.io.parse_example(
      eg[tf.newaxis], {
          'image/encoded': tf.io.FixedLenFeature(shape=(), dtype=tf.string),
          'image/text': tf.io.FixedLenFeature(shape=(), dtype=tf.string)
      })
  return example['image/encoded'][0], example['image/text'][0]
img, txt = tf_parse(raw_example)
print(txt.numpy())
print(repr(img.numpy()[:20]), "...")
b'Rue Perreyon'
b'\x89PNG\r\n\x1a\n\x00\x00\x00\rIHDR\x00\x00\x02X' ...
decoded = dataset.map(tf_parse)
decoded
<_MapDataset element_spec=(TensorSpec(shape=(), dtype=tf.string, name=None), TensorSpec(shape=(), dtype=tf.string, name=None))>
image_batch, text_batch = next(iter(decoded.batch(10)))
image_batch.shape
TensorShape([10])

時間序列視窗化

有關端到端時間序列示例,請參閱:時間序列預測

時間序列資料通常是在保持時間軸完整的情況下組織的。

使用簡單的 Dataset.range 來演示

range_ds = tf.data.Dataset.range(100000)

通常,基於此類資料的模型將需要連續的時間切片。

最簡單的方法是對資料進行批處理

使用 batch

batches = range_ds.batch(10, drop_remainder=True)

for batch in batches.take(5):
  print(batch.numpy())
[0 1 2 3 4 5 6 7 8 9]
[10 11 12 13 14 15 16 17 18 19]
[20 21 22 23 24 25 26 27 28 29]
[30 31 32 33 34 35 36 37 38 39]
[40 41 42 43 44 45 46 47 48 49]

或者為了向未來一步進行密集預測,您可以將特徵和標籤相對於彼此移動一步

def dense_1_step(batch):
  # Shift features and labels one step relative to each other.
  return batch[:-1], batch[1:]

predict_dense_1_step = batches.map(dense_1_step)

for features, label in predict_dense_1_step.take(3):
  print(features.numpy(), " => ", label.numpy())
[0 1 2 3 4 5 6 7 8]  =>  [1 2 3 4 5 6 7 8 9]
[10 11 12 13 14 15 16 17 18]  =>  [11 12 13 14 15 16 17 18 19]
[20 21 22 23 24 25 26 27 28]  =>  [21 22 23 24 25 26 27 28 29]

要預測整個視窗而不是固定偏移量,您可以將批次拆分為兩部分

batches = range_ds.batch(15, drop_remainder=True)

def label_next_5_steps(batch):
  return (batch[:-5],   # Inputs: All except the last 5 steps
          batch[-5:])   # Labels: The last 5 steps

predict_5_steps = batches.map(label_next_5_steps)

for features, label in predict_5_steps.take(3):
  print(features.numpy(), " => ", label.numpy())
[0 1 2 3 4 5 6 7 8 9]  =>  [10 11 12 13 14]
[15 16 17 18 19 20 21 22 23 24]  =>  [25 26 27 28 29]
[30 31 32 33 34 35 36 37 38 39]  =>  [40 41 42 43 44]

為了允許一個批次的特徵和另一個批次的標籤之間存在一些重疊,請使用 Dataset.zip

feature_length = 10
label_length = 3

features = range_ds.batch(feature_length, drop_remainder=True)
labels = range_ds.batch(feature_length).skip(1).map(lambda labels: labels[:label_length])

predicted_steps = tf.data.Dataset.zip((features, labels))

for features, label in predicted_steps.take(5):
  print(features.numpy(), " => ", label.numpy())
[0 1 2 3 4 5 6 7 8 9]  =>  [10 11 12]
[10 11 12 13 14 15 16 17 18 19]  =>  [20 21 22]
[20 21 22 23 24 25 26 27 28 29]  =>  [30 31 32]
[30 31 32 33 34 35 36 37 38 39]  =>  [40 41 42]
[40 41 42 43 44 45 46 47 48 49]  =>  [50 51 52]

使用 window

雖然使用 Dataset.batch 有效,但在某些情況下您可能需要更精細的控制。Dataset.window 方法為您提供完全控制,但需要小心:它返回一個 DatasetsDataset。有關詳細資訊,請轉到資料集結構部分。

window_size = 5

windows = range_ds.window(window_size, shift=1)
for sub_ds in windows.take(5):
  print(sub_ds)
<_VariantDataset element_spec=TensorSpec(shape=(), dtype=tf.int64, name=None)>
<_VariantDataset element_spec=TensorSpec(shape=(), dtype=tf.int64, name=None)>
<_VariantDataset element_spec=TensorSpec(shape=(), dtype=tf.int64, name=None)>
<_VariantDataset element_spec=TensorSpec(shape=(), dtype=tf.int64, name=None)>
<_VariantDataset element_spec=TensorSpec(shape=(), dtype=tf.int64, name=None)>

Dataset.flat_map 方法可以接收資料集的資料集並將它們扁平化為單個數據集

for x in windows.flat_map(lambda x: x).take(30):
   print(x.numpy(), end=' ')
0 1 2 3 4 1 2 3 4 5 2 3 4 5 6 3 4 5 6 7 4 5 6 7 8 5 6 7 8 9

在幾乎所有情況下,您都希望首先 Dataset.batch 該資料集

def sub_to_batch(sub):
  return sub.batch(window_size, drop_remainder=True)

for example in windows.flat_map(sub_to_batch).take(5):
  print(example.numpy())
[0 1 2 3 4]
[1 2 3 4 5]
[2 3 4 5 6]
[3 4 5 6 7]
[4 5 6 7 8]

現在,您可以看到 shift 引數控制每個視窗移動的距離。

綜合起來,您可以編寫此函式

def make_window_dataset(ds, window_size=5, shift=1, stride=1):
  windows = ds.window(window_size, shift=shift, stride=stride)

  def sub_to_batch(sub):
    return sub.batch(window_size, drop_remainder=True)

  windows = windows.flat_map(sub_to_batch)
  return windows
ds = make_window_dataset(range_ds, window_size=10, shift = 5, stride=3)

for example in ds.take(10):
  print(example.numpy())
[ 0  3  6  9 12 15 18 21 24 27]
[ 5  8 11 14 17 20 23 26 29 32]
[10 13 16 19 22 25 28 31 34 37]
[15 18 21 24 27 30 33 36 39 42]
[20 23 26 29 32 35 38 41 44 47]
[25 28 31 34 37 40 43 46 49 52]
[30 33 36 39 42 45 48 51 54 57]
[35 38 41 44 47 50 53 56 59 62]
[40 43 46 49 52 55 58 61 64 67]
[45 48 51 54 57 60 63 66 69 72]

然後像以前一樣輕鬆提取標籤

dense_labels_ds = ds.map(dense_1_step)

for inputs,labels in dense_labels_ds.take(3):
  print(inputs.numpy(), "=>", labels.numpy())
[ 0  3  6  9 12 15 18 21 24] => [ 3  6  9 12 15 18 21 24 27]
[ 5  8 11 14 17 20 23 26 29] => [ 8 11 14 17 20 23 26 29 32]
[10 13 16 19 22 25 28 31 34] => [13 16 19 22 25 28 31 34 37]

重取樣

在使用類不平衡嚴重的資料集時,您可能需要重取樣該資料集。tf.data 提供了兩種方法來執行此操作。信用卡欺詐資料集是此類問題的一個很好的例子。

zip_path = tf.keras.utils.get_file(
    origin='https://storage.googleapis.com/download.tensorflow.org/data/creditcard.zip',
    fname='creditcard.zip',
    extract=True)

csv_path = zip_path.replace('.zip', '.csv')
Downloading data from https://storage.googleapis.com/download.tensorflow.org/data/creditcard.zip
69155632/69155632 ━━━━━━━━━━━━━━━━━━━━ 1s 0us/step
creditcard_ds = tf.data.experimental.make_csv_dataset(
    csv_path, batch_size=1024, label_name="Class",
    # Set the column types: 30 floats and an int.
    column_defaults=[float()]*30+[int()])

現在,檢查類別的分佈,它高度傾斜

def count(counts, batch):
  features, labels = batch
  class_1 = labels == 1
  class_1 = tf.cast(class_1, tf.int32)

  class_0 = labels == 0
  class_0 = tf.cast(class_0, tf.int32)

  counts['class_0'] += tf.reduce_sum(class_0)
  counts['class_1'] += tf.reduce_sum(class_1)

  return counts
counts = creditcard_ds.take(10).reduce(
    initial_state={'class_0': 0, 'class_1': 0},
    reduce_func = count)

counts = np.array([counts['class_0'].numpy(),
                   counts['class_1'].numpy()]).astype(np.float32)

fractions = counts/counts.sum()
print(fractions)
[0.996 0.004]

使用不平衡資料集進行訓練的一種常見方法是平衡它。tf.data 包括幾種實現此工作流的方法

資料集取樣

重取樣資料集的一種方法是使用 sample_from_datasets。這在您為每個類別都有一個單獨的 tf.data.Dataset 時更適用。

在這裡,只需使用 filter 從信用卡欺詐資料中生成它們

negative_ds = (
  creditcard_ds
    .unbatch()
    .filter(lambda features, label: label==0)
    .repeat())
positive_ds = (
  creditcard_ds
    .unbatch()
    .filter(lambda features, label: label==1)
    .repeat())
for features, label in positive_ds.batch(10).take(1):
  print(label.numpy())
[1 1 1 1 1 1 1 1 1 1]

要使用 tf.data.Dataset.sample_from_datasets,請傳遞資料集以及每個資料集的權重

balanced_ds = tf.data.Dataset.sample_from_datasets(
    [negative_ds, positive_ds], [0.5, 0.5]).batch(10)

現在資料集以 50/50 的機率產生每個類別的示例

for features, labels in balanced_ds.take(10):
  print(labels.numpy())
[1 0 1 0 0 1 0 1 1 0]
[1 0 0 0 0 0 0 0 1 1]
[0 0 1 0 0 1 0 0 1 0]
[0 1 1 0 1 0 0 1 1 0]
[0 1 1 0 0 0 1 1 1 1]
[1 1 1 1 1 1 0 0 0 0]
[0 1 1 0 1 0 0 1 1 1]
[1 1 0 0 0 0 0 1 0 1]
[1 1 0 1 1 1 1 0 0 1]
[0 1 1 0 0 1 0 0 0 0]

拒絕重取樣

上述 Dataset.sample_from_datasets 方法的一個問題是它需要每個類別一個單獨的 tf.data.Dataset。您可以使用 Dataset.filter 來建立這兩個資料集,但這會導致所有資料被載入兩次。

tf.data.Dataset.rejection_resample 方法可以應用於資料集以重新平衡它,同時只加載一次。元素將被丟棄或重複以實現平衡。

rejection_resample 方法接受一個 class_func 引數。此 class_func 被應用於每個資料集元素,並用於確定為了平衡目的,示例屬於哪個類別。

這裡的目標是平衡標籤分佈,而 creditcard_ds 的元素已經是 (features, label) 對。因此 class_func 只需返回這些標籤

def class_func(features, label):
  return label

重取樣方法處理單個示例,因此在這種情況下,您必須在應用該方法之前 unbatch 資料集。

該方法需要目標分佈,並可選擇初始分佈估計作為輸入。

resample_ds = (
    creditcard_ds
    .unbatch()
    .rejection_resample(class_func, target_dist=[0.5,0.5],
                        initial_dist=fractions)
    .batch(10))
WARNING:tensorflow:From /tmpfs/src/tf_docs_env/lib/python3.9/site-packages/tensorflow/python/data/ops/dataset_ops.py:4968: Print (from tensorflow.python.ops.logging_ops) is deprecated and will be removed after 2018-08-20.
Instructions for updating:
Use tf.print instead of tf.Print. Note that tf.print returns a no-output operator that directly prints the output. Outside of defuns or eager mode, this operator will not be executed unless it is directly specified in session.run or used as a control dependency for other operators. This is only a concern in graph mode. Below is an example of how to ensure tf.print executes in graph mode:

rejection_resample 方法返回 (class, example) 對,其中 classclass_func 的輸出。在這種情況下,example 已經是一個 (feature, label) 對,因此使用 map 丟棄標籤的額外副本

balanced_ds = resample_ds.map(lambda extra_label, features_and_label: features_and_label)

現在資料集以 50/50 的機率產生每個類別的示例

for features, labels in balanced_ds.take(10):
  print(labels.numpy())
Proportion of examples rejected by sampler is high: [0.995996118][0.995996118 0.00400390616][0 1]
Proportion of examples rejected by sampler is high: [0.995996118][0.995996118 0.00400390616][0 1]
Proportion of examples rejected by sampler is high: [0.995996118][0.995996118 0.00400390616][0 1]
Proportion of examples rejected by sampler is high: [0.995996118][0.995996118 0.00400390616][0 1]
Proportion of examples rejected by sampler is high: [0.995996118][0.995996118 0.00400390616][0 1]
Proportion of examples rejected by sampler is high: [0.995996118][0.995996118 0.00400390616][0 1]
Proportion of examples rejected by sampler is high: [0.995996118][0.995996118 0.00400390616][0 1]
Proportion of examples rejected by sampler is high: [0.995996118][0.995996118 0.00400390616][0 1]
Proportion of examples rejected by sampler is high: [0.995996118][0.995996118 0.00400390616][0 1]
Proportion of examples rejected by sampler is high: [0.995996118][0.995996118 0.00400390616][0 1]
[1 0 1 0 1 0 1 0 1 1]
[1 0 1 1 1 1 0 0 1 0]
[1 0 1 1 0 1 0 0 0 1]
[0 1 0 0 0 0 1 1 1 1]
[1 0 0 0 1 1 1 0 1 0]
[0 0 0 1 0 0 1 0 1 1]
[0 1 0 0 0 0 1 0 1 0]
[1 0 0 0 0 1 0 0 0 1]
[0 0 0 0 1 1 1 1 1 0]
[1 1 0 1 1 1 1 1 1 0]

迭代器檢查點

Tensorflow 支援進行檢查點設定,以便當您的訓練過程重新啟動時,它可以恢復最新的檢查點以恢復其大部分進度。除了檢查點模型變數外,您還可以檢查資料集迭代器的進度。如果您擁有一個大型資料集,並且不想在每次重新啟動時從頭開始資料集,這可能很有用。但請注意,迭代器檢查點可能很大,因為諸如 Dataset.shuffleDataset.prefetch 之類的轉換需要在迭代器內緩衝元素。

要將您的迭代器包含在檢查點中,請將迭代器傳遞給 tf.train.Checkpoint 建構函式。

range_ds = tf.data.Dataset.range(20)

iterator = iter(range_ds)
ckpt = tf.train.Checkpoint(step=tf.Variable(0), iterator=iterator)
manager = tf.train.CheckpointManager(ckpt, '/tmp/my_ckpt', max_to_keep=3)

print([next(iterator).numpy() for _ in range(5)])

save_path = manager.save()

print([next(iterator).numpy() for _ in range(5)])

ckpt.restore(manager.latest_checkpoint)

print([next(iterator).numpy() for _ in range(5)])
[0, 1, 2, 3, 4]
[5, 6, 7, 8, 9]
[5, 6, 7, 8, 9]

tf.datatf.keras 一起使用

tf.keras API 簡化了建立和執行機器學習模型的許多方面。其 Model.fitModel.evaluateModel.predict API 支援資料集作為輸入。這是一個快速的資料集和模型設定

train, test = tf.keras.datasets.fashion_mnist.load_data()

images, labels = train
images = images/255.0
labels = labels.astype(np.int32)
fmnist_train_ds = tf.data.Dataset.from_tensor_slices((images, labels))
fmnist_train_ds = fmnist_train_ds.shuffle(5000).batch(32)

model = tf.keras.Sequential([
  tf.keras.layers.Flatten(),
  tf.keras.layers.Dense(10)
])

model.compile(optimizer='adam',
              loss=tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True),
              metrics=['accuracy'])

傳遞一個 (feature, label) 對資料集對於 Model.fitModel.evaluate 來說就足夠了

model.fit(fmnist_train_ds, epochs=2)
Epoch 1/2
WARNING: All log messages before absl::InitializeLog() is called are written to STDERR
I0000 00:00:1723685884.693688   45100 service.cc:146] XLA service 0x7f35cc006690 initialized for platform CUDA (this does not guarantee that XLA will be used). Devices:
I0000 00:00:1723685884.693721   45100 service.cc:154]   StreamExecutor device (0): Tesla T4, Compute Capability 7.5
I0000 00:00:1723685884.693725   45100 service.cc:154]   StreamExecutor device (1): Tesla T4, Compute Capability 7.5
I0000 00:00:1723685884.693728   45100 service.cc:154]   StreamExecutor device (2): Tesla T4, Compute Capability 7.5
I0000 00:00:1723685884.693731   45100 service.cc:154]   StreamExecutor device (3): Tesla T4, Compute Capability 7.5
136/1875 ━━━━━━━━━━━━━━━━━━━━ 1s 1ms/step - accuracy: 0.5241 - loss: 1.4346
I0000 00:00:1723685885.241810   45100 device_compiler.h:188] Compiled cluster using XLA!  This line is logged at most once for the lifetime of the process.
1875/1875 ━━━━━━━━━━━━━━━━━━━━ 3s 1ms/step - accuracy: 0.7449 - loss: 0.7643
Epoch 2/2
1875/1875 ━━━━━━━━━━━━━━━━━━━━ 2s 1ms/step - accuracy: 0.8381 - loss: 0.4704
<keras.src.callbacks.history.History at 0x7f373f583250>

如果您傳遞一個無限資料集,例如透過呼叫 Dataset.repeat,您只需要同時傳遞 steps_per_epoch 引數

model.fit(fmnist_train_ds.repeat(), epochs=2, steps_per_epoch=20)
Epoch 1/2
20/20 ━━━━━━━━━━━━━━━━━━━━ 0s 1ms/step - accuracy: 0.8254 - loss: 0.4682  
Epoch 2/2
20/20 ━━━━━━━━━━━━━━━━━━━━ 0s 1ms/step - accuracy: 0.8622 - loss: 0.4263
<keras.src.callbacks.history.History at 0x7f37443e1190>

對於評估,您可以傳遞評估步驟數

loss, accuracy = model.evaluate(fmnist_train_ds)
print("Loss :", loss)
print("Accuracy :", accuracy)
1875/1875 ━━━━━━━━━━━━━━━━━━━━ 2s 1ms/step - accuracy: 0.8504 - loss: 0.4343
Loss : 0.4353208839893341
Accuracy : 0.849216639995575

對於長資料集,設定要評估的步驟數

loss, accuracy = model.evaluate(fmnist_train_ds.repeat(), steps=10)
print("Loss :", loss)
print("Accuracy :", accuracy)
10/10 ━━━━━━━━━━━━━━━━━━━━ 0s 2ms/step - accuracy: 0.8411 - loss: 0.5209  
Loss : 0.46679750084877014
Accuracy : 0.84375

呼叫 Model.predict 時不需要標籤。

predict_ds = tf.data.Dataset.from_tensor_slices(images).batch(32)
result = model.predict(predict_ds, steps = 10)
print(result.shape)
10/10 ━━━━━━━━━━━━━━━━━━━━ 0s 1ms/step  
(320, 10)

但如果您確實傳遞了包含標籤的資料集,標籤將被忽略

result = model.predict(fmnist_train_ds, steps = 10)
print(result.shape)
10/10 ━━━━━━━━━━━━━━━━━━━━ 0s 1ms/step 
(320, 10)