mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-10 12:21:57 +08:00
fix bugs in semantic segmentation example (#1824)
* Update unet.py * Update semantic_segmentation.py
This commit is contained in:
@@ -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}
|
||||
|
||||
|
||||
@@ -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))
|
||||
|
||||
Reference in New Issue
Block a user