(lhs, rhs, dout)
| 36 | |
| 37 | @jit.xla_trace(without_host=True) |
| 38 | def func(lhs, rhs, dout): |
| 39 | gm.attach([lhs, rhs]) |
| 40 | with gm: |
| 41 | out = F.matmul(lhs, rhs, lhs_transpose, rhs_transpose) |
| 42 | gm.backward(out, dout) |
| 43 | return out, lhs.grad, rhs.grad |
| 44 | |
| 45 | mge_rsts = func(lhs, rhs, dout) |
| 46 | xla_rsts = func(lhs, rhs, dout) |
no test coverage detected