MCPcopy Create free account
hub / github.com/zai-org/CodeGeeX / get_op_support_info

Function get_op_support_info

codegeex/mindspore/scripts/layer_norm.py:43–92  ·  view source on GitHub ↗

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")

Source from the content-addressed store, hash-verified

41# 'pylint: disable = unused-argument
42# 'pylint: disable=too-many-arguments,too-many-locals
43def 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

Callers

nothing calls this directly

Calls 2

to_frac_z_axisFunction · 0.85
getMethod · 0.45

Tested by

no test coverage detected