fix bugs in semantic segmentation example (#1824)

* Update unet.py

* Update semantic_segmentation.py
This commit is contained in:
Nand Dalal
2020-05-14 02:36:45 -04:00
committed by GitHub
parent 1265b2fe02
commit cf2d32d0a6
2 changed files with 2 additions and 2 deletions
@@ -165,7 +165,7 @@ class SegModel(pl.LightningModule):
return {'val_loss': loss_val}
def validation_epoch_end(self, outputs):
loss_val = sum(output['val_loss'] for output in outputs) / len(outputs)
loss_val = torch.stack([x['val_loss'] for x in outputs]).mean()
log_dict = {'val_loss': loss_val}
return {'log': log_dict, 'val_loss': log_dict['val_loss'], 'progress_bar': log_dict}
+1 -1
View File
@@ -33,7 +33,7 @@ class UNet(nn.Module):
feats *= 2
for _ in range(num_layers - 1):
layers.append(Up(feats, feats // 2), bilinear)
layers.append(Up(feats, feats // 2, bilinear))
feats //= 2
layers.append(nn.Conv2d(feats, num_classes, kernel_size=1))