MCPcopy Create free account
hub / github.com/apache/tvm / test_logaddexp

Function test_logaddexp

tests/python/relax/test_frontend_from_exported_program.py:974–1016  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

972
973
974def 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
1019def test_logical_and():

Callers

nothing calls this directly

Calls 2

LogAddExpClass · 0.85
verify_modelFunction · 0.70

Tested by

no test coverage detected

Used in the wild real call sites across dependent graphs

searching dependent graphs…