()
| 972 | |
| 973 | |
| 974 | def test_logaddexp(): |
| 975 | class LogAddExp(Module): |
| 976 | def forward(self, input1, input2): |
| 977 | return torch.logaddexp(input1, input2) |
| 978 | |
| 979 | @tvm.script.ir_module |
| 980 | class expected: |
| 981 | @R.function |
| 982 | def main( |
| 983 | input1: R.Tensor((1, 3, 10, 10), dtype="float32"), |
| 984 | input2: R.Tensor((1, 3, 10, 10), dtype="float32"), |
| 985 | ) -> R.Tuple(R.Tensor((1, 3, 10, 10), dtype="float32")): |
| 986 | # block 0 |
| 987 | with R.dataflow(): |
| 988 | lv: R.Tensor((1, 3, 10, 10), dtype="bool") = R.greater_equal(input1, input2) |
| 989 | lv1: R.Tensor((1, 3, 10, 10), dtype="float32") = R.where(lv, input1, input2) |
| 990 | lv2: R.Tensor((1, 3, 10, 10), dtype="float32") = R.where(lv, input2, input1) |
| 991 | lv3: R.Tensor((1, 3, 10, 10), dtype="float32") = R.abs(input1) |
| 992 | lv4: R.Tensor((1, 3, 10, 10), dtype="bool") = R.not_equal( |
| 993 | lv3, R.const(float("inf"), "float32") |
| 994 | ) |
| 995 | lv5: R.Tensor((1, 3, 10, 10), dtype="bool") = R.equal(input1, input1) |
| 996 | lv6: R.Tensor((1, 3, 10, 10), dtype="bool") = R.multiply(lv5, lv4) |
| 997 | lv7: R.Tensor((1, 3, 10, 10), dtype="bool") = R.logical_not(lv6) |
| 998 | lv8: R.Tensor((1, 3, 10, 10), dtype="bool") = R.equal(input1, input2) |
| 999 | lv9: R.Tensor((1, 3, 10, 10), dtype="bool") = R.logical_and(lv7, lv8) |
| 1000 | lv10: R.Tensor((1, 3, 10, 10), dtype="float32") = R.subtract(lv2, lv1) |
| 1001 | lv11: R.Tensor((1, 3, 10, 10), dtype="float32") = R.exp(lv10) |
| 1002 | lv12: R.Tensor((1, 3, 10, 10), dtype="float32") = R.add( |
| 1003 | lv11, R.const(1.0, "float32") |
| 1004 | ) |
| 1005 | lv13: R.Tensor((1, 3, 10, 10), dtype="float32") = R.log(lv12) |
| 1006 | lv14: R.Tensor((1, 3, 10, 10), dtype="float32") = R.add(lv1, lv13) |
| 1007 | lv15: R.Tensor((1, 3, 10, 10), dtype="float32") = R.where(lv9, input1, lv14) |
| 1008 | gv: R.Tuple(R.Tensor((1, 3, 10, 10), dtype="float32")) = (lv15,) |
| 1009 | R.output(gv) |
| 1010 | return gv |
| 1011 | |
| 1012 | example_args = ( |
| 1013 | torch.randn(1, 3, 10, 10, dtype=torch.float32), |
| 1014 | torch.randn(1, 3, 10, 10, dtype=torch.float32), |
| 1015 | ) |
| 1016 | verify_model(LogAddExp(), example_args, {}, expected) |
| 1017 | |
| 1018 | |
| 1019 | def test_logical_and(): |
nothing calls this directly
no test coverage detected
searching dependent graphs…