(in_layout, axes, keep_dims)
| 111 | |
| 112 | |
| 113 | def get_expected_layout(in_layout, axes, keep_dims): |
| 114 | in_layout = in_layout or "" |
| 115 | if keep_dims or not in_layout: |
| 116 | return in_layout |
| 117 | if axes is None: |
| 118 | return "" |
| 119 | if isinstance(axes, int): |
| 120 | axes = [axes] |
| 121 | ndim = len(in_layout) |
| 122 | axes = [(axis + ndim) % ndim for axis in axes] |
| 123 | return "".join(c for i, c in enumerate(in_layout) if i not in axes) |
| 124 | |
| 125 | |
| 126 | def check_layout(tensor, in_layout, axes, keep_dims): |