(
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,
use_program_cache=False,
)
| 2643 | return None |
| 2644 | |
| 2645 | def _prepare_pipeline_ctx( |
| 2646 | self, |
| 2647 | program=None, |
| 2648 | dataset=None, |
| 2649 | scope=None, |
| 2650 | thread=0, |
| 2651 | is_infer=False, |
| 2652 | debug=False, |
| 2653 | fetch_list=None, |
| 2654 | fetch_info=None, |
| 2655 | print_period=100, |
| 2656 | fetch_handler=None, |
| 2657 | use_program_cache=False, |
| 2658 | ): |
| 2659 | assert program._pipeline_opt is not None |
| 2660 | assert dataset is None, "dataset should be None for pipeline mode" |
| 2661 | |
| 2662 | cache_key = _get_strong_program_cache_key(program, None, fetch_list) |
| 2663 | ctx = self._get_ctx_cache(cache_key) |
| 2664 | if use_program_cache and ctx is not None: |
| 2665 | return ctx |
| 2666 | |
| 2667 | import paddle |
| 2668 | |
| 2669 | # The following fake dataset is created to call |
| 2670 | # the _prepare_trainer api, and it is meaningless. |
| 2671 | def _get_dataset(): |
| 2672 | data_vars = [] |
| 2673 | for var in program.global_block().vars.values(): |
| 2674 | if var.is_data: |
| 2675 | data_vars.append(var) |
| 2676 | dataset = paddle.base.DatasetFactory().create_dataset( |
| 2677 | 'FileInstantDataset' |
| 2678 | ) |
| 2679 | dataset.set_batch_size(1) |
| 2680 | dataset.set_thread(1) |
| 2681 | dataset.set_filelist(['None']) |
| 2682 | dataset.set_use_var(data_vars) |
| 2683 | dataset._prepare_to_run() |
| 2684 | return dataset |
| 2685 | |
| 2686 | dataset = _get_dataset() |
| 2687 | |
| 2688 | def _get_real_program_fetch_list(): |
| 2689 | real_program = program._pipeline_opt["section_program"] |
| 2690 | real_fetch_list = [] |
| 2691 | for fetch_var in fetch_list: |
| 2692 | if isinstance(fetch_var, Variable): |
| 2693 | fetch_var_name = fetch_var.name |
| 2694 | else: |
| 2695 | fetch_var_name = fetch_var |
| 2696 | if fetch_var_name in real_program.global_block().vars: |
| 2697 | real_fetch_list.append(fetch_var) |
| 2698 | |
| 2699 | real_program = _add_feed_fetch_ops( |
| 2700 | program=real_program, |
| 2701 | feed=[], |
| 2702 | fetch_list=real_fetch_list, |
no test coverage detected