| 358 | |
| 359 | |
| 360 | class 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 | |
| 400 | class DataFeeder: |
no outgoing calls
no test coverage detected