MCPcopy Create free account
hub / github.com/PaddlePaddle/Paddle / _prepare_pipeline_ctx

Method _prepare_pipeline_ctx

python/paddle/base/executor.py:2645–2750  ·  view source on GitHub ↗
(
        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,
    )

Source from the content-addressed store, hash-verified

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,

Callers 1

_run_pipelineMethod · 0.95

Calls 8

_get_ctx_cacheMethod · 0.95
_prepare_trainerMethod · 0.95
_add_ctx_cacheMethod · 0.95
_set_inferMethod · 0.45
_gen_trainer_descMethod · 0.45
_descMethod · 0.45

Tested by

no test coverage detected