创建PyReader用于Paddle读取数据 Args: args: 模型参数,定义于utils.DefaultArgs file_name: string类型,数据文件路径 feed_list: list类型,模型输入的列表 place: Paddle执行的空间,即GPU和CPU reader: 读取数据用的类,定义与reader.py iterable: 是否返回可迭代的PyReader for_test: 是否用于测试,如果测试则不shuffle Retur
(args, file_name, feed_list, place,
reader=None, iterable=True, for_test=False)
| 195 | |
| 196 | |
| 197 | def create_pyreader(args, file_name, feed_list, place, |
| 198 | reader=None, iterable=True, for_test=False): |
| 199 | """创建PyReader用于Paddle读取数据 |
| 200 | |
| 201 | Args: |
| 202 | args: 模型参数,定义于utils.DefaultArgs |
| 203 | file_name: string类型,数据文件路径 |
| 204 | feed_list: list类型,模型输入的列表 |
| 205 | place: Paddle执行的空间,即GPU和CPU |
| 206 | reader: 读取数据用的类,定义与reader.py |
| 207 | iterable: 是否返回可迭代的PyReader |
| 208 | for_test: 是否用于测试,如果测试则不shuffle |
| 209 | |
| 210 | Returns: |
| 211 | PyReader对象,用于迭代读取数据 |
| 212 | """ |
| 213 | # init reader |
| 214 | pyreader = fluid.io.PyReader( |
| 215 | feed_list=feed_list, |
| 216 | capacity=50, |
| 217 | use_double_buffer=True, |
| 218 | iterable=iterable |
| 219 | ) |
| 220 | if reader is None: |
| 221 | reader = Dataset(args) |
| 222 | |
| 223 | if for_test: |
| 224 | pyreader.decorate_sample_list_generator( |
| 225 | paddle.batch( |
| 226 | reader.file_reader(file_name, mode='test'), |
| 227 | batch_size=args.batch_size |
| 228 | ), |
| 229 | places=place |
| 230 | ) |
| 231 | else: |
| 232 | pyreader.decorate_sample_list_generator( |
| 233 | paddle.batch( |
| 234 | paddle.reader.shuffle( |
| 235 | reader.file_reader(file_name), |
| 236 | buf_size=args.traindata_shuffle_buffer |
| 237 | ), |
| 238 | batch_size=args.batch_size |
| 239 | ), |
| 240 | places=place |
| 241 | ) |
| 242 | |
| 243 | return pyreader |
| 244 | |
| 245 | |
| 246 | def test_process(exe, program, reader, test_ret): |
no test coverage detected