| 16 | |
| 17 | |
| 18 | class KernelInfo: |
| 19 | def __init__(self, op_type): |
| 20 | self.op_type = op_type |
| 21 | self.supported_dtypes = set() |
| 22 | |
| 23 | def parse_phi_dtypes(self, registered_info_list, device="GPU"): |
| 24 | assert isinstance(registered_info_list, list) |
| 25 | assert device in ["CPU", "GPU"] |
| 26 | |
| 27 | # registered_info_list is in format as follows: |
| 28 | # ['(GPU, Undefined(AnyLayout), float32)', '(GPU, Undefined(AnyLayout), float64)'] |
| 29 | for kernel_str in registered_info_list: |
| 30 | kernel_strs = ( |
| 31 | kernel_str.replace("(", "").replace(")", "").split(",") |
| 32 | ) |
| 33 | # for GPU, the kernel type can be GPUDNN. |
| 34 | if device in kernel_strs[0]: |
| 35 | self.supported_dtypes.add(kernel_strs[-1].replace(" ", "")) |
| 36 | |
| 37 | # if len(self.supported_dtypes) == 0: |
| 38 | # print("-- [WARNING] No dtypes for op_type={}, device={}. Registered info: {}".format(self.op_type, device, registered_info_list)) |
| 39 | |
| 40 | def parse_fluid_dtypes(self, registered_info_list, device="gpu"): |
| 41 | assert isinstance(registered_info_list, list) |
| 42 | assert device in ["cpu", "gpu"] |
| 43 | |
| 44 | # registered_info_list is in format as follows: |
| 45 | # ['{data_type[::paddle::platform::bfloat16]; data_layout[Undefined(AnyLayout)]; place[Place(gpu:0)]; library_type[PLAIN]}', ...}'] |
| 46 | for kernel_str in registered_info_list: |
| 47 | kernel_strs = kernel_str.split(";") |
| 48 | if "place" in kernel_strs[2] and device in kernel_strs[2]: |
| 49 | assert "data_type" in kernel_strs[0] |
| 50 | dtype_str = kernel_strs[0].replace("{data_type[", "") |
| 51 | dtype_str = dtype_str.replace("::paddle::platform::", "") |
| 52 | dtype_str = dtype_str.replace("]", "") |
| 53 | self.supported_dtypes.add(dtype_str) |
| 54 | |
| 55 | |
| 56 | class KernelRegistryStatistics: |
no outgoing calls
no test coverage detected