| 237 | |
| 238 | |
| 239 | def test_shape_pattern(): |
| 240 | shape = [32, 32] |
| 241 | pattern = wildcard().has_shape(shape) |
| 242 | assert isinstance(pattern, ShapePattern) |
| 243 | tvm_ffi.structural_equal(pattern.shape, shape) |
| 244 | assert pattern.match(bindings[0].var) |
| 245 | assert wildcard().has_shape([32, 32]).match(bindings[0].var) |
| 246 | n, m = tirx.Var("n", dtype="int64"), tirx.Var("m", dtype="int64") |
| 247 | symsh_var = rx.Var("x", R.Tensor([n, m, n + m], "float32")) |
| 248 | assert wildcard().has_shape([n, m, n + m]).match(symsh_var) |
| 249 | assert wildcard().has_shape([n, m, m + n]).match(symsh_var) # + is commutative. |
| 250 | assert not wildcard().has_shape([1, 2, 3]).match(symsh_var) |
| 251 | assert not wildcard().has_shape([m, n, n + m]).match(symsh_var) |
| 252 | |
| 253 | |
| 254 | def test_prim_arr_pattern(): |