| 103 | |
| 104 | |
| 105 | def 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 | |
| 132 | def test_var_binding() -> None: |