(
self,
program=None,
dataset=None,
scope=None,
thread=0,
is_infer=False,
debug=False,
fetch_list=None,
fetch_info=None,
print_period=100,
fetch_handler=None,
)
| 2476 | return scope, trainer |
| 2477 | |
| 2478 | def _run_from_dataset( |
| 2479 | self, |
| 2480 | program=None, |
| 2481 | dataset=None, |
| 2482 | scope=None, |
| 2483 | thread=0, |
| 2484 | is_infer=False, |
| 2485 | debug=False, |
| 2486 | fetch_list=None, |
| 2487 | fetch_info=None, |
| 2488 | print_period=100, |
| 2489 | fetch_handler=None, |
| 2490 | ): |
| 2491 | if program._pipeline_opt is not None: |
| 2492 | import paddle |
| 2493 | |
| 2494 | if dataset is not None: |
| 2495 | raise RuntimeError("dataset should be None for pipeline mode") |
| 2496 | # The following fake dataset is created to call |
| 2497 | # the _prepare_trainer api, and it is meaningless. |
| 2498 | data_vars = [] |
| 2499 | for var in program.global_block().vars.values(): |
| 2500 | if var.is_data: |
| 2501 | data_vars.append(var) |
| 2502 | dataset = paddle.base.DatasetFactory().create_dataset( |
| 2503 | 'FileInstantDataset' |
| 2504 | ) |
| 2505 | dataset.set_batch_size(1) |
| 2506 | dataset.set_thread(1) |
| 2507 | dataset.set_filelist(['None']) |
| 2508 | dataset.set_use_var(data_vars) |
| 2509 | elif program._heter_pipeline_opt is not None: |
| 2510 | stage_id = program._heter_pipeline_opt["pipeline_stage"] |
| 2511 | # print("test_fl_stage_id: {}".format(stage_id)) |
| 2512 | heter_place = program._heter_pipeline_opt["heter_place"] |
| 2513 | if stage_id != 0: |
| 2514 | if "is_fl_mode" not in program._heter_pipeline_opt: |
| 2515 | import paddle |
| 2516 | |
| 2517 | if dataset is not None: |
| 2518 | raise RuntimeError( |
| 2519 | "dataset should be None for heter pipeline mode" |
| 2520 | ) |
| 2521 | # The following fake dataset is created to call |
| 2522 | # the _prepare_trainer api, and it is meaningless. |
| 2523 | data_vars = [] |
| 2524 | for var in program.global_block().vars.values(): |
| 2525 | if var.is_data: |
| 2526 | data_vars.append(var) |
| 2527 | dataset = paddle.base.DatasetFactory().create_dataset( |
| 2528 | 'InMemoryDataset' |
| 2529 | ) |
| 2530 | dataset.set_batch_size(1) |
| 2531 | dataset.set_thread(1) |
| 2532 | dataset.set_filelist(['None']) |
| 2533 | dataset.set_use_var(data_vars) |
| 2534 | else: |
| 2535 | if dataset is None: |
no test coverage detected