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

Function test_shape_pattern

tests/python/relax/test_dataflow_pattern.py:239–251  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

237
238
239def 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
254def test_prim_arr_pattern():

Callers

nothing calls this directly

Calls 4

wildcardFunction · 0.85
has_shapeMethod · 0.80
TensorMethod · 0.80
matchMethod · 0.45

Tested by

no test coverage detected

Used in the wild real call sites across dependent graphs

searching dependent graphs…