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

Class BatchedTensorProvider

python/paddle/base/data_feeder.py:360–397  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

358
359
360class BatchedTensorProvider:
361 def __init__(self, feed_list, place, batch_size, generator, drop_last):
362 self.place = place
363 self.batch_size = batch_size
364 self.generator = generator
365 self.converters = []
366 self.drop_last = drop_last
367
368 for var in feed_list:
369 if not in_pir_mode():
370 assert var.lod_level == 0, "lod_level must be 0"
371 self.converters.append(
372 DataToDenseTensorConverter(
373 place=self.place,
374 lod_level=0,
375 shape=var.shape,
376 dtype=var.dtype,
377 )
378 )
379
380 def _done(self):
381 return [c.done() for c in self.converters]
382
383 def __call__(self):
384 idx = 0
385 for each_sample in self.generator():
386 for each_slot, each_converter in zip(each_sample, self.converters):
387 each_converter.data.append(each_slot)
388
389 idx += 1
390 if idx == self.batch_size:
391 idx = 0
392 yield self._done()
393
394 if not self.drop_last and idx > 0:
395 yield self._done()
396 else:
397 [c._reset() for c in self.converters]
398
399
400class DataFeeder:

Callers 1

set_sample_generatorMethod · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected