tf.data.experimental.dense_to_ragged_batch
Преобразование, которое группирует разрозненные элементы в tf.RaggedTensor.
tf.data.experimental.dense_to_ragged_batch(
batch_size,
drop_remainder=False,
row_splits_dtype=tf.dtypes.int64
)
Это преобразование объединяет несколько последовательных элементов входного набора данных в один элемент.
Как и в случае с tf.data.Dataset.batch, компоненты результирующего элемента будут иметь дополнительное внешнее измерение, которое будет batch_size (или N % batch_size для последнего элемента, если batch_size не делит количество входных элементов N равномерно, и drop_remainder является False). Если ваша программа зависит от того, что у партий одинаковое внешнее измерение, вы должны установить аргумент drop_remainder в True, чтобы предотвратить создание меньшей партии.
В отличие от tf.data.Dataset.batch, входные элементы, подлежащие объединению в пакеты, могут иметь разные формы:
- Если входной элемент является
tf.Tensor, чья статическаяtf.TensorShapeполностью определена, то он группируется в пакеты как обычно. - Если входной элемент является
tf.Tensor, чья статическаяtf.TensorShapeсодержит одну или несколько осей с неизвестным размером (т. е.shape[i]=None), то выходной элемент будет содержатьtf.RaggedTensor, которая разрознена по любой из таких размерностей. - Если входной элемент является
tf.RaggedTensorили любого другого типа, то он группируется в пакеты как обычно.
Пример:
dataset = tf.data.Dataset.from_tensor_slices(np.arange(6))
dataset = dataset.map(lambda x: tf.range(x))
dataset.element_spec.shape
TensorShape([None])
dataset = dataset.apply(
tf.data.experimental.dense_to_ragged_batch(batch_size=2))
for batch in dataset:
print(batch)
<tf.RaggedTensor [[], [0]]>
<tf.RaggedTensor [[0, 1], [0, 1, 2]]>
<tf.RaggedTensor [[0, 1, 2, 3], [0, 1, 2, 3, 4]]>
| Аргументы | |
|---|---|
batch_size | Скалярное значение tf.int64 tf.Tensor, представляющее количество последовательных элементов этого набора данных, объединяемых в одну партию. |
drop_remainder | (Необязательно.) Скалярное значение tf.bool tf.Tensor, представляющее, должна ли быть пропущена последняя партия в случае, если в ней меньше batch_size элементов; поведение по умолчанию заключается в том, чтобы не пропускать меньшую партию. |
row_splits_dtype | Тип данных, который должен использоваться для row_splits любых новых разрозненных тензоров. Существующие элементы tf.RaggedTensor не изменяют свой тип данных row_splits. |
| Возвращаемое значение | |
|---|---|
Dataset | Dataset. |
© 2022 The TensorFlow Authors. All rights reserved.
Licensed under the Creative Commons Attribution License 4.0.
Code samples licensed under the Apache 2.0 License.
https://www.tensorflow.org/versions/r2.9/api_docs/python/tf/data/experimental/dense_to_ragged_batch