(inps: Iterable[Tensor])
| 47 | |
| 48 | |
| 49 | def _infer_broadcasted_shape(inps: Iterable[Tensor]) -> tuple: |
| 50 | broadcasted_ndim = inps[0].ndim |
| 51 | broadcasted_shape = list(inps[0]._tuple_shape) |
| 52 | for i in range(1, len(inps)): |
| 53 | cur_ndim = inps[i].ndim |
| 54 | cur_shape = list(inps[i]._tuple_shape) |
| 55 | n_dim = max(cur_ndim, broadcasted_ndim) |
| 56 | for j in range(n_dim - 1, -1, -1): |
| 57 | cur_dim = cur_ndim + j - n_dim |
| 58 | broad_dim = broadcasted_ndim + j - n_dim |
| 59 | cur_size = cur_shape[cur_dim] if cur_dim >= 0 else 1 |
| 60 | broad_size = broadcasted_shape[broad_dim] if broad_dim >= 0 else 1 |
| 61 | assert cur_size == broad_size or cur_size == 1 or broad_size == 1, ( |
| 62 | "The size of inps[{}] ({}) must match the size ({}) at " |
| 63 | "dim {}".format(i, cur_size, broad_size, j) |
| 64 | ) |
| 65 | broad_size = max(cur_size, broad_size) |
| 66 | if broad_dim < 0: |
| 67 | broadcasted_shape = [broad_size] + broadcasted_shape |
| 68 | broadcasted_ndim += 1 |
| 69 | else: |
| 70 | broadcasted_shape[broad_dim] = broad_size |
| 71 | return tuple(broadcasted_shape) |
| 72 | |
| 73 | |
| 74 | def _broadcast_tensors_with_size( |
no test coverage detected