()
| 37 | |
| 38 | |
| 39 | def test_allreduce_infer_struct_info(): |
| 40 | bb = relax.BlockBuilder() |
| 41 | x0 = relax.Var("x", R.Tensor((2, 3), "float32")) |
| 42 | x1 = relax.Var("x", R.Tensor("float32", ndim=3)) |
| 43 | x2 = relax.Var("x", R.Tensor("float32", ndim=-1)) |
| 44 | x3 = relax.Var("x", R.Tensor((2, 3))) |
| 45 | x4 = relax.Var("x", R.Tensor()) |
| 46 | x5 = relax.Var("x", R.Tensor((3, 4))) |
| 47 | |
| 48 | _check_inference(bb, relax.op.ccl.allreduce(x0), relax.TensorStructInfo((2, 3), "float32")) |
| 49 | _check_inference( |
| 50 | bb, relax.op.ccl.allreduce(x1), relax.TensorStructInfo(dtype="float32", ndim=3) |
| 51 | ) |
| 52 | _check_inference(bb, relax.op.ccl.allreduce(x2), relax.TensorStructInfo(dtype="float32")) |
| 53 | _check_inference(bb, relax.op.ccl.allreduce(x3), relax.TensorStructInfo((2, 3), dtype="")) |
| 54 | _check_inference(bb, relax.op.ccl.allreduce(x4), relax.TensorStructInfo(dtype="")) |
| 55 | _check_inference(bb, relax.op.ccl.allreduce(x5), relax.TensorStructInfo((3, 4), dtype="")) |
| 56 | |
| 57 | |
| 58 | def test_allreduce_infer_struct_info_shape_symbolic(): |
nothing calls this directly
no test coverage detected
searching dependent graphs…