MCPcopy Create free account
hub / github.com/PaddlePaddle/Paddle / decorate_batch_generator

Method decorate_batch_generator

python/paddle/base/reader.py:1568–1637  ·  view source on GitHub ↗

Set the data source of the PyReader object. The provided :code:`reader` should be a Python generator, which yields numpy.ndarray-typed or DenseTensor-typed batched data. :code:`places` must be set when the PyReader object is iterable. Args: rea

(self, reader, places=None)

Source from the content-addressed store, hash-verified

1566 self._loader.set_sample_list_generator(reader, places)
1567
1568 def decorate_batch_generator(self, reader, places=None):
1569 '''
1570 Set the data source of the PyReader object.
1571
1572 The provided :code:`reader` should be a Python generator,
1573 which yields numpy.ndarray-typed or DenseTensor-typed batched data.
1574
1575 :code:`places` must be set when the PyReader object is iterable.
1576
1577 Args:
1578 reader (generator): Python generator that yields DenseTensor-typed
1579 batched data.
1580 places (None|list(CUDAPlace)|list(CPUPlace)): place list. Must
1581 be provided when PyReader is iterable.
1582
1583 Example:
1584 .. code-block:: pycon
1585
1586 >>> import paddle
1587 >>> import paddle.base as base
1588 >>> import numpy as np
1589
1590 >>> paddle.enable_static()
1591
1592 >>> EPOCH_NUM = 3
1593 >>> ITER_NUM = 15
1594 >>> BATCH_SIZE = 3
1595
1596 >>> def network(image, label):
1597 ... # User-defined network, here is an example of softmax regression.
1598 ... predict = paddle.static.nn.fc(x=image, size=10, activation='softmax')
1599 ... return paddle.nn.functional.cross_entropy(
1600 ... input=predict,
1601 ... label=label,
1602 ... reduction='none',
1603 ... use_softmax=False,
1604 ... )
1605
1606 >>> def random_image_and_label_generator(height, width):
1607 ... def generator():
1608 ... for i in range(ITER_NUM):
1609 ... batch_image = np.random.uniform(low=0, high=255, size=[BATCH_SIZE, height, width])
1610 ... batch_label = np.ones([BATCH_SIZE, 1])
1611 ... batch_image = batch_image.astype('float32')
1612 ... batch_label = batch_label.astype('int64')
1613 ... yield batch_image, batch_label
1614 ...
1615 ... return generator
1616
1617 >>> image = paddle.static.data(name='image', shape=[None, 784, 784], dtype='float32')
1618 >>> label = paddle.static.data(name='label', shape=[None, 1], dtype='int64')
1619 >>> reader = base.io.PyReader(feed_list=[image, label], capacity=4, iterable=True)
1620
1621 >>> user_defined_generator = random_image_and_label_generator(784, 784)
1622 >>> reader.decorate_batch_generator(
1623 ... user_defined_generator,
1624 ... paddle.CPUPlace(),
1625 ... )

Callers 1

main_implMethod · 0.95

Calls 1

set_batch_generatorMethod · 0.45

Tested by 1

main_implMethod · 0.76