Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
15 changes: 10 additions & 5 deletions monai/losses/perceptual.py
Original file line number Diff line number Diff line change
Expand Up @@ -49,22 +49,27 @@ class PerceptualLoss(nn.Module):

Args:
spatial_dims: number of spatial dimensions.
network_type: {``"alex"``, ``"vgg"``, ``"squeeze"``, ``"radimagenet_resnet50"``,
``"medicalnet_resnet10_23datasets"``, ``"medicalnet_resnet50_23datasets"``, ``"resnet50"``}
Specifies the network architecture to use. Defaults to ``"alex"``.
network_type: type of network for perceptual loss. One of:
- "alex"
- "vgg"
- "squeeze"
- "radimagenet_resnet50"
- "medicalnet_resnet10_23datasets"
- "medicalnet_resnet50_23datasets"
- "resnet50"
is_fake_3d: if True use 2.5D approach for a 3D perceptual loss.
fake_3d_ratio: ratio of how many slices per axis are used in the 2.5D approach.
cache_dir: path to cache directory to save the pretrained network weights.
pretrained: whether to load pretrained weights. This argument only works when using networks from
LIPIS or Torchvision. Defaults to ``"True"``.
LIPIS or Torchvision. Defaults to ``True``.
pretrained_path: if `pretrained` is `True`, users can specify a weights file to be loaded
via using this argument. This argument only works when ``"network_type"`` is "resnet50".
Defaults to `None`.
pretrained_state_dict_key: if `pretrained_path` is not `None`, this argument is used to
extract the expected state dict. This argument only works when ``"network_type"`` is "resnet50".
Defaults to `None`.
channel_wise: if True, the loss is returned per channel. Otherwise the loss is averaged over the channels.
Defaults to ``False``.
Defaults to ``False``.
"""

def __init__(
Expand Down
Loading