(inp, axis)
| 79 | |
| 80 | |
| 81 | def expand_dims(inp, axis): |
| 82 | assert isinstance(axis, int), f"only int axis supported, get {axis}" |
| 83 | assert ( |
| 84 | axis >= -inp.ndim - 1 and axis <= inp.ndim |
| 85 | ), f"invalid axis {axis} for {inp.shape}" |
| 86 | |
| 87 | dst_shape = list(inp.shape) |
| 88 | insert_pos = axis if axis >= 0 else (axis + inp.ndim + 1) |
| 89 | dst_shape.insert(insert_pos, 1) |
| 90 | |
| 91 | return inp.reshape(tuple(dst_shape)) |
| 92 | |
| 93 | |
| 94 | @register_lower_rule(mops.Dimshuffle) |
no test coverage detected