From cf2d32d0a6c757aad39c36b621a646ed3a24619a Mon Sep 17 00:00:00 2001 From: Nand Dalal Date: Thu, 14 May 2020 01:36:45 -0500 Subject: [PATCH] fix bugs in semantic segmentation example (#1824) * Update unet.py * Update semantic_segmentation.py --- pl_examples/domain_templates/semantic_segmentation.py | 2 +- pl_examples/models/unet.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/pl_examples/domain_templates/semantic_segmentation.py b/pl_examples/domain_templates/semantic_segmentation.py index 4604b645..9d98c799 100644 --- a/pl_examples/domain_templates/semantic_segmentation.py +++ b/pl_examples/domain_templates/semantic_segmentation.py @@ -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} diff --git a/pl_examples/models/unet.py b/pl_examples/models/unet.py index 5e85802b..36fc6573 100644 --- a/pl_examples/models/unet.py +++ b/pl_examples/models/unet.py @@ -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))