(register_te_grads)
| 119 | |
| 120 | |
| 121 | def test_emit_te(register_te_grads): |
| 122 | # Build the target module using emit_te |
| 123 | def f_mul(src1, src2): |
| 124 | def mul(*idx): |
| 125 | return src1[idx] * src2[idx] |
| 126 | |
| 127 | return tvm.te.compute(src1.shape, mul, name="f_mul") |
| 128 | |
| 129 | a = relax.Var("a", relax.TensorStructInfo([5, 5], "float32")) |
| 130 | b = relax.Var("b", relax.TensorStructInfo([5, 5], "float32")) |
| 131 | |
| 132 | bb = relax.BlockBuilder() |
| 133 | with bb.function("main", [a, b]): |
| 134 | with bb.dataflow(): |
| 135 | d = bb.emit( |
| 136 | bb.call_te_with_grad( |
| 137 | f_mul, a, b, primfunc_name_hint="f_mul", te_grad_name="f_mul_grad" |
| 138 | ) |
| 139 | ) |
| 140 | out = bb.emit_output(R.sum(d)) |
| 141 | bb.emit_func_output(out) |
| 142 | |
| 143 | Before = bb.get() |
| 144 | After = Gradient("main")(Before) |
| 145 | assert_structural_equal(After, get_expected_1()) |
| 146 | |
| 147 | |
| 148 | def test_call_tir(register_te_grads): |
nothing calls this directly
no test coverage detected
searching dependent graphs…