Function
train_one_step
(model, criterion, optimizer, inputs, targets)
Source from the content-addressed store, hash-verified
| 121 | |
| 122 | |
| 123 | def train_one_step(model, criterion, optimizer, inputs, targets): |
| 124 | # convert to tensors |
| 125 | inputs = torch.from_numpy(inputs.astype(np.float32)) |
| 126 | targets = torch.from_numpy(targets.astype(np.float32)) |
| 127 | |
| 128 | # zero the parameter gradients |
| 129 | optimizer.zero_grad() |
| 130 | |
| 131 | # Forward pass |
| 132 | outputs = model(inputs) |
| 133 | loss = criterion(outputs, targets) |
| 134 | |
| 135 | # Backward and optimize |
| 136 | loss.backward() |
| 137 | optimizer.step() |
| 138 | |
| 139 | |
| 140 | |
Tested by
no test coverage detected