common_shape = [1] * reduce(max, (len(shape) for shape in shapes)) for arg_idx, shape in enumerate(shapes): for idx in range(-1, -1 - len(shape), -1): # align from the right if common_shape[idx] == 1: common_shape[idx] = shape[idx] # 1 stretches to anything if shape[idx] == 1: continue # ... in either direction torch._check(common_shape[idx] == shape[idx], # otherwise sizes must match lambda: f"Attempting to broadcast a dimension of length {shape[idx]} at {idx}!") return common_shape