MCPcopy Create free account
hub / github.com/MegEngine/MegEngine / _infer_broadcasted_shape

Function _infer_broadcasted_shape

imperative/python/megengine/random/rng.py:49–71  ·  view source on GitHub ↗
(inps: Iterable[Tensor])

Source from the content-addressed store, hash-verified

47
48
49def _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
74def _broadcast_tensors_with_size(

Callers 1

Calls 3

listFunction · 0.85
maxFunction · 0.85
formatMethod · 0.45

Tested by

no test coverage detected