| 1042 | self._reader.reset() |
| 1043 | |
| 1044 | def set_sample_generator( |
| 1045 | self, reader, batch_size, drop_last=True, places=None |
| 1046 | ): |
| 1047 | assert batch_size > 0, "batch_size must be larger than 0" |
| 1048 | if isinstance(places, (list, tuple)): |
| 1049 | places = _get_paddle_place_list(places) |
| 1050 | else: |
| 1051 | places = _get_paddle_place(places) |
| 1052 | has_lod = False |
| 1053 | if not in_pir_mode(): |
| 1054 | for f in self._feed_list: |
| 1055 | if f.lod_level != 0: |
| 1056 | has_lod = True |
| 1057 | break |
| 1058 | |
| 1059 | if has_lod: |
| 1060 | self.set_sample_list_generator( |
| 1061 | paddle.batch( |
| 1062 | reader, batch_size=batch_size, drop_last=drop_last |
| 1063 | ), |
| 1064 | places=places, |
| 1065 | ) |
| 1066 | else: |
| 1067 | reader = BatchedTensorProvider( |
| 1068 | feed_list=self._feed_list, |
| 1069 | place=core.CPUPlace(), |
| 1070 | batch_size=batch_size, |
| 1071 | generator=reader, |
| 1072 | drop_last=drop_last, |
| 1073 | ) |
| 1074 | self.set_batch_generator(reader, places=places) |
| 1075 | return self |
| 1076 | |
| 1077 | def set_sample_list_generator(self, reader, places=None): |
| 1078 | if isinstance(places, (list, tuple)): |