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

Function test_match_cast

tests/python/relax/test_ast_printer.py:105–129  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

103
104
105def test_match_cast() -> None:
106 # match_cast([16, 8], [m, n])
107 m = tirx.Var("m", dtype="int64")
108 n = tirx.Var("n", dtype="int64")
109 shape = rx.const([16, 8], "int32")
110 var = rx.Var("v0", R.Shape())
111 b0 = rx.MatchCast(var, shape, R.Tensor([m, n], "int32"))
112 b0_str = dump_ast(b0)
113 assert b0_str.startswith("MatchCast(")
114 assert "Constant" in b0_str
115 assert "PrimExpr(value=`m" in b0_str
116 assert "PrimExpr(value=`n" in b0_str
117 assert "16" in b0_str
118 assert "8" in b0_str
119
120 # var1: Tensor((m, n), "float32") =
121 # match_cast(var0: R.Tensor("float32"), [m, n])
122 value = rx.Var("value", R.Tensor("float32"))
123 var = rx.Var("v1", R.Tensor([m, n], "float32"))
124 b1 = rx.MatchCast(var, value, R.Tensor([m, n], "float32"))
125 b1_str = dump_ast(b1)
126 assert b1_str.startswith("MatchCast(")
127 assert "PrimExpr(value=`m" in b1_str
128 assert "PrimExpr(value=`n" in b1_str
129 assert b1_str != dump_ast(b1, include_struct_info_annotations=False)
130
131
132def test_var_binding() -> None:

Callers

nothing calls this directly

Calls 3

dump_astFunction · 0.90
ShapeMethod · 0.80
TensorMethod · 0.80

Tested by

no test coverage detected

Used in the wild real call sites across dependent graphs

searching dependent graphs…