U tPh4@szddlmZddlZddlmZddlmZmZddlm Z dddd d Z dddd d Z e e d Z GdddeZ dS)) annotationsN)_Loss) gaussian_1dseparable_filtering) LossReductionint torch.Tensor)sigmareturncCs,|dkrtd|tt|ddddS)Nr$expecting positive sigma, got sigma=sampledF)r truncatedapprox normalize) ValueErrorrtorchtensorr rM/home/dell461/cl/sdc2/HISourceFinder-master-l/src/monai/losses/multi_scale.pymake_gaussian_kernelsrcsbdkrtdtd}tfddt| |dD}t|}|t|}|S)Nrr csg|]}|ddqS)r).0xrrr sz&make_cauchy_kernel..r)rrrrrange reciprocalsum)r tailkrrrmake_cauchy_kernels $ r#)gaussiancauchycsJeZdZdZddejfdddddd fd d Zd d d d ddZZS)MultiScaleLossz This is a wrapper class. It smooths the input and target at different scales before passing them into the wrapped loss function. Adapted from: DeepReg (https://github.com/DeepRegNet/DeepReg) Nr$rz list | NonestrzLossReduction | strNone)lossscaleskernel reductionr csFtjt|jd|tkr,td|dt||_||_||_dS)z Args: loss: loss function to be wrapped scales: list of scalars or None, if None, do not apply any scaling. kernel: gaussian or cauchy. )r,zgot unsupported kernel type: z only support gaussian and cauchyN) super__init__rvaluekernel_fn_dictr kernel_fnr)r*)selfr)r*r+r, __class__rrr.1s  zMultiScaleLoss.__init__r)y_truey_predr c Cs|jdkr|||}ng}|jD]n}|dkrB||||q"||t||||g|jdt||||g|jdq"tj|dd}|j t j j krt |}n:|j t jj krt|}n |j t jj krtd|j d|S)Nrr)dimzUnsupported reduction: z0, available options are ["mean", "sum", "none"].)r*r)appendrr1tondimrstackr,rMEANr/meanSUMr NONEr)r2r5r6r) loss_listsrrrforwardEs(      zMultiScaleLoss.forward) __name__ __module__ __qualname____doc__rr<r.rB __classcell__rrr3rr&(s  r&) __future__rrtorch.nn.modules.lossrmonai.networks.layersrr monai.utilsrrr#r0r&rrrr s