(self, api_item_yaml)
| 49 | |
| 50 | class BaseAPI: |
| 51 | def __init__(self, api_item_yaml): |
| 52 | self.api = self.get_api_name(api_item_yaml) |
| 53 | |
| 54 | # inputs: |
| 55 | # names : [], list of input names |
| 56 | # input_info : {input_name : type} |
| 57 | # attrs: |
| 58 | # names : [], list of attribute names |
| 59 | # attr_info : { attr_name : (type, default_values)} |
| 60 | # outputs: |
| 61 | # names : [], list of output names |
| 62 | # types : [], list of output types |
| 63 | # out_size_expr : [], expression for getting size of vector<Tensor> |
| 64 | ( |
| 65 | self.inputs, |
| 66 | self.attrs, |
| 67 | self.outputs, |
| 68 | self.optional_vars, |
| 69 | ) = self.parse_args(self.api, api_item_yaml) |
| 70 | |
| 71 | self.is_base_api = True |
| 72 | self.is_only_composite_api = False |
| 73 | # Whether to generate code for inplace API |
| 74 | self.is_inplace_context = False |
| 75 | if 'invoke' in api_item_yaml: |
| 76 | self.is_base_api = False |
| 77 | self.invoke = api_item_yaml['invoke'] |
| 78 | else: |
| 79 | if 'infer_meta' in api_item_yaml: |
| 80 | self.infer_meta = self.parse_infer_meta( |
| 81 | api_item_yaml['infer_meta'] |
| 82 | ) |
| 83 | if 'composite' in api_item_yaml and 'kernel' not in api_item_yaml: |
| 84 | self.is_base_api = False |
| 85 | self.is_only_composite_api = True |
| 86 | self.kernel = None |
| 87 | else: |
| 88 | self.kernel = self.parse_kernel(api_item_yaml['kernel']) |
| 89 | self.data_transform = self.parse_data_transform(api_item_yaml) |
| 90 | self.inplace_map, self.view_map = {}, {} |
| 91 | |
| 92 | self.gene_input_func = { |
| 93 | "const Tensor&": { |
| 94 | "dense": self.gene_dense_input, |
| 95 | "selected_rows": self.gene_selected_rows_input, |
| 96 | }, |
| 97 | "const paddle::optional<Tensor>&": { |
| 98 | "dense": self.gene_dense_input, |
| 99 | "selected_rows": self.gene_selected_rows_input, |
| 100 | }, |
| 101 | "const std::vector<Tensor>&": {"dense": self.gene_vec_dense_input}, |
| 102 | "const paddle::optional<std::vector<Tensor>>&": { |
| 103 | "dense": self.gene_optional_vec_dense_input |
| 104 | }, |
| 105 | } |
| 106 | |
| 107 | def get_api_name(self, api_item_yaml): |
| 108 | if 'op' in api_item_yaml: |
nothing calls this directly
no test coverage detected