| 191 | |
| 192 | |
| 193 | def train(model, train_loader, optimizer, epoch, device, verbose = 0, |
| 194 | lossFn = None, lr_schedule=None, |
| 195 | post_proc = lambda args: args): |
| 196 | |
| 197 | if lossFn is None: |
| 198 | lossFn = nn.MSELoss() |
| 199 | |
| 200 | model.train() |
| 201 | |
| 202 | total_loss = 0. |
| 203 | |
| 204 | for batch_idx, (data, target) in enumerate(train_loader): |
| 205 | |
| 206 | bs = len(data) |
| 207 | data, target = data.to(device), target.to(device) |
| 208 | optimizer.zero_grad() |
| 209 | |
| 210 | output = model(data) |
| 211 | |
| 212 | target = post_proc(target) |
| 213 | output = post_proc(output) |
| 214 | loss = lossFn(output.view(bs, -1), target.view(bs, -1)) |
| 215 | |
| 216 | loss.backward() |
| 217 | optimizer.step() |
| 218 | total_loss += loss.sum().item() |
| 219 | if lr_schedule is not None: lr_schedule.step() |
| 220 | |
| 221 | if verbose>0: |
| 222 | print('Train Epoch: {} [{}/{} ({:.0f}%)]\tLoss: {:.6f}'.format( |
| 223 | epoch, batch_idx * len(data), len(train_loader.dataset), |
| 224 | 100. * batch_idx / len(train_loader), loss.item())) |
| 225 | |
| 226 | return total_loss/len(train_loader.dataset) |
| 227 | |
| 228 | |
| 229 | def test(model, test_loader, device, verbose=0, lossFn=None, |