(d)
| 37 | data = np.random.random((6, nchannel, 2, 2)).astype("float32") |
| 38 | |
| 39 | def train_fn(d): |
| 40 | for _ in range(3): |
| 41 | with gm: |
| 42 | loss = m(d).mean() |
| 43 | gm.backward(loss) |
| 44 | optim.step() |
| 45 | return loss |
| 46 | |
| 47 | if use_trace: |
| 48 | train_fn = trace(train_fn, symbolic=use_symbolic) |
no test coverage detected