U tPh@sbddlmZddlmZddlZddlmmZddl m Z ddl m Z Gddde Z e ZdS)) annotations)UnionN)_Loss) pytorch_aftercsdeZdZdZddddddfd d Zdd d dddZddddddZddddddZZS)DeepSupervisionLossz Wrapper class around the main loss function to accept a list of tensors returned from a deeply supervised networks. The final loss is computed as the sum of weighted losses for each of deep supervision levels. expNrstrzlist[float] | NoneNone)loss weight_modeweightsreturncs4t||_||_||_tddr*dnd|_dS)a Args: loss: main loss instance, e.g DiceLoss(). weight_mode: {``"same"``, ``"exp"``, ``"two"``} Specifies the weights calculation for each image level. Defaults to ``"exp"``. - ``"same"``: all weights are equal to 1. - ``"exp"``: exponentially decreasing weights by a power of 2: 1, 0.5, 0.25, 0.125, etc . - ``"two"``: equal smaller weights for lower levels: 1, 0.5, 0.5, 0.5, 0.5, etc weights: a list of weights to apply to each deeply supervised sub-loss, if provided, this will be used regardless of the weight_mode  z nearest-exactnearestN)super__init__r r r r interp_mode)selfr r r  __class__I/home/dell461/cl/sdc2/HISourceFinder-master-l/src/monai/losses/ds_loss.pyrs zDeepSupervisionLoss.__init__rintz list[float])levelsr cCstd|}|jdk r2t|j|kr2|jd|}n\|jdkrHdg|}nF|jdkrfddt|D}n(|jdkrd dt|D}n dg|}|S) zG Calculates weights for a given number of scale levels rNsame?rcSsg|]}td|dqS)?g?)max.0lrrr 9sz3DeepSupervisionLoss.get_weights..twocSsg|]}|dkrdndqS)rrrrrrrrr";s)rr lenr range)rrr rrr get_weights/s      zDeepSupervisionLoss.get_weightsz torch.Tensor)inputtargetr cCsD|jdd|jddkr8tj||jdd|jd}|||S)z Calculates a loss output accounting for differences in shapes, and downsizing targets if necessary (using nearest neighbor interpolation) Generally downsizing occurs for all level, except for the first (level==0) N)sizemode)shapeF interpolaterr )rr'r(rrrget_lossAszDeepSupervisionLoss.get_lossz-Union[None, torch.Tensor, list[torch.Tensor]]cCst|ttfrh|jt|d}tjdtj|jd}t t|D]$}|||| |||7}q>|S|dkrxt d| ||S)N)rr)dtypedevicezinput shouldn't be None.) isinstancelisttupler&r$torchtensorfloatr1r%r/ ValueErrorr )rr'r(r r r!rrrforwardKs"zDeepSupervisionLoss.forward)rN)r) __name__ __module__ __qualname____doc__rr&r/r9 __classcell__rrrrrs  r) __future__rtypingrr5torch.nn.functionalnn functionalr-torch.nn.modules.lossr monai.utilsrrds_lossrrrr s    A