(iters)
| 262 | check_trace=False) |
| 263 | |
| 264 | def train(iters): |
| 265 | for _ in range(iters): |
| 266 | # Get some fake data |
| 267 | inp = torch.randn(5, 1, 28, 28, device='cuda') |
| 268 | out = traced_net(inp) |
| 269 | |
| 270 | # Here's some fake loss |
| 271 | out.sum().backward() |
| 272 | |
| 273 | # Zero out grads |
| 274 | traced_net.zero_grad() |
| 275 | |
| 276 | # Set it up so the params have .grad fields so they are not reported as leaks |
| 277 | train(1) |
no test coverage detected