get_op_support_info
(input_x, input_gamma, input_beta,
output_y, output_mean, output_variance,
begin_norm_axis, begin_params_axis,
epsilon=1e-12, kernel_name="layer_norm",
impl_mode="high_performance")
| 41 | # 'pylint: disable = unused-argument |
| 42 | # 'pylint: disable=too-many-arguments,too-many-locals |
| 43 | def get_op_support_info(input_x, input_gamma, input_beta, |
| 44 | output_y, output_mean, output_variance, |
| 45 | begin_norm_axis, begin_params_axis, |
| 46 | epsilon=1e-12, kernel_name="layer_norm", |
| 47 | impl_mode="high_performance"): |
| 48 | """ |
| 49 | get_op_support_info |
| 50 | """ |
| 51 | format_x = input_x.get("format").upper() |
| 52 | shape_x = input_x.get("shape") |
| 53 | ori_shape_x = input_x.get("ori_shape") |
| 54 | begin_norm_axis = shape_util.axis_check(len(shape_x), begin_norm_axis) |
| 55 | begin_params_axis = shape_util.axis_check(len(shape_x), begin_params_axis) |
| 56 | axis_split_matrix = [] |
| 57 | |
| 58 | if format_x in ("ND", "NCHW", "NHWC", "NC1HWC0"): |
| 59 | if begin_params_axis == 0: |
| 60 | for i in range(begin_norm_axis): |
| 61 | split_0 = [SplitInput([0, [i], [-1], [-1]], [1, [i], [-1], [-1]], [2, [i], [-1], [-1]]), |
| 62 | SplitOutput([0, [i]], [1, [i]], [2, [i]])] |
| 63 | axis_split_matrix.append(split_0) |
| 64 | else: |
| 65 | if begin_norm_axis <= begin_params_axis: |
| 66 | for i in range(begin_norm_axis): |
| 67 | split_0 = [SplitInput([0, [i], [-1], [-1]]), |
| 68 | SplitOutput([0, [i]], [1, [i]], [2, [i]])] |
| 69 | axis_split_matrix.append(split_0) |
| 70 | else: |
| 71 | for i in range(begin_params_axis): |
| 72 | split_0 = [SplitInput([0, [i], [-1], [-1]]), |
| 73 | SplitOutput([0, [i]], [1, [i]], [2, [i]])] |
| 74 | axis_split_matrix.append(split_0) |
| 75 | |
| 76 | elif format_x == "FRACTAL_NZ": |
| 77 | index_list = tuple(index for index, _ in enumerate(ori_shape_x)) |
| 78 | start_axis = min(begin_norm_axis, begin_params_axis) |
| 79 | |
| 80 | no_split_axis = index_list[start_axis:] |
| 81 | no_split_axis = to_frac_z_axis(ori_shape_x, no_split_axis) |
| 82 | for i in range(len(shape_x)): |
| 83 | if i not in no_split_axis: |
| 84 | split_0 = [SplitInput([0, [i], [-1], [-1]]), |
| 85 | SplitOutput([0, [i]], [1, [i]], [2, [i]])] |
| 86 | axis_split_matrix.append(split_0) |
| 87 | |
| 88 | else: |
| 89 | axis_split_matrix = None |
| 90 | axis_reduce_list = None |
| 91 | op_cal_info_in_json = get_op_cal_info(axis_split_matrix, axis_reduce_list, 0, 0) |
| 92 | return op_cal_info_in_json |
| 93 | |
| 94 | |
| 95 | # 'pylint: disable=locally-disabled,too-many-arguments,unused-argument |
nothing calls this directly
no test coverage detected