(x, w, b, dy)
| 35 | |
| 36 | @jit.xla_trace(without_host=True) |
| 37 | def func(x, w, b, dy): |
| 38 | gm.attach([x, w, b]) |
| 39 | with gm: |
| 40 | y = F.conv2d( |
| 41 | x, |
| 42 | w, |
| 43 | b, |
| 44 | stride=stride, |
| 45 | padding=padding, |
| 46 | groups=groups, |
| 47 | compute_mode=cm, |
| 48 | ) |
| 49 | gm.backward(y, dy) |
| 50 | return [y, x.grad, w.grad, b.grad] |
| 51 | |
| 52 | mge_rsts = func(x, w, b, dy) |
| 53 | xla_rsts = func(x, w, b, dy) |