o , i&@sZddlmZddlmZddlZddlmZddlmZed\Z Z Gdddej Z dS)) annotations)castN)optional_importztorchvision.modelscsFeZdZdZ      d d!fdd Zd"ddZd#d$ddZZS)%MILModela Multiple Instance Learning (MIL) model, with a backbone classification model. Currently, it only works for 2D images, a typical use case is for classification of the digital pathology whole slide images. The expected shape of input data is `[B, N, C, H, W]`, where `B` is the batch_size of PyTorch Dataloader and `N` is the number of instances extracted from every original image in the batch. A tutorial example is available at: https://github.com/Project-MONAI/tutorials/tree/master/pathology/multiple_instance_learning. Args: num_classes: number of output classes. mil_mode: MIL algorithm, available values (Defaults to ``"att"``): - ``"mean"`` - average features from all instances, equivalent to pure CNN (non MIL). - ``"max"`` - retain only the instance with the max probability for loss calculation. - ``"att"`` - attention based MIL https://arxiv.org/abs/1802.04712. - ``"att_trans"`` - transformer MIL https://arxiv.org/abs/2111.01556. - ``"att_trans_pyramid"`` - transformer pyramid MIL https://arxiv.org/abs/2111.01556. pretrained: init backbone with pretrained weights, defaults to ``True``. backbone: Backbone classifier CNN (either ``None``, a ``nn.Module`` that returns features, or a string name of a torchvision model). Defaults to ``None``, in which case ResNet50 is used. backbone_num_features: Number of output features of the backbone CNN Defaults to ``None`` (necessary only when using a custom backbone) trans_blocks: number of the blocks in `TransformEncoder` layer. trans_dropout: dropout rate in `TransformEncoder` layer. attTN num_classesintmil_modestr pretrainedboolbackbonestr | nn.Module | Nonebackbone_num_features int | None trans_blocks trans_dropoutfloatreturnNonec s.t|dkrtdt||dvrtdt||_t_d_ |durtt j |r8t j j ndd}|jj} tj|_i_|dkrsfdd} |j| d |j| d |j| d |j| d nSt|trtt |d} | durtd t|| |rdndd}t|dddur|jj} tj|_n tdt|dt|tjr|}|} |durtdntd|dur|dvrtdt|jdvrnjdkrtt| dttdd_njdkrtj| d|d} tj| |d_ tt| dttdd_nmjdkrttjtjdd|d|dttddtjtjdd|d|dttd dtjtjdd|d|dtjtjd!d|d|dg} | _ | d} tt| dttdd_ntdt|t| |_ |_!dS)"Nrz$Number of classes must be positive: )meanmaxr att_transatt_trans_pyramidzUnsupported mil_mode: )weightsrcsfdd}|S)Ncs|j<dS)N) extra_outputs)moduleinputoutput) layer_nameself^/home/dell461/cl/sdc2/last_ska_mid/HISourceFinder-master-l/src/monai/networks/nets/milmodel.pyhookVsz5MILModel.__init__..forward_hook..hookr#)r!r%r")r!r$ forward_hookTsz'MILModel.__init__..forward_hooklayer1layer2layer3layer4zUnknown torch vision modelDEFAULTfcz4Unable to detect FC layer for the torchvision model z0. Please initialize the backbone model manually.zJNumber of endencoder features must be provided for a custom backbone modelzUnsupported backbone)rrrrz.Custom backbone is not supported for the mode:)rrrir)d_modelnheaddropout) num_layersiii )"super__init__ ValueErrorr lowerr nn Sequential attention transformermodelsresnet50ResNet50_Weights IMAGENET1K_V1r- in_featurestorchIdentityrr(register_forward_hookr)r*r+ isinstancegetattrModuleLinearTanhTransformerEncoderLayerTransformerEncoder ModuleListmyfcnet)r"r r r rrrrrNnfcr'Z torch_modelr<transformer_list __class__r&r$r65s            & &   & zMILModel.__init__x torch.Tensorc Cs|j}|jdkr||}tj|dd}|S|jdkr+||}tj|dd\}}|S|jdkrL||}tj|dd}tj||dd}||}|S|jdkr|j dur| ddd}| |}| ddd}||}tj|dd}tj||dd}||}|S|jd krH|j durHtj|j d d d |d|dd  ddd}tj|j d d d |d|dd  ddd}tj|j dd d |d|dd  ddd}tj|j dd d |d|dd  ddd}t tj|j } | d|}| dtj||fdd}| dtj||fdd}| dtj||fdd}| ddd}||}tj|dd}tj||dd}||}|Stdt|j)Nrr.)dimrrrrrr()rVr)r*r+rWzWrong model mode)shaper rMrBrrr;softmaxsumr<permuterreshaperr9rLcatr7r ) r"rSsh_al1l2l3l4rPr#r#r$ calc_headsR  0 ,  %   0000   zMILModel.calc_headFno_headcCs`|j}||d|d|d|d|d}||}||d|dd}|s.||}|S)Nrr.rVrWrrX)rYr]rNrf)r"rSrgr_r#r#r$forwards(  zMILModel.forward)rTNNrr)r r r r r rrrrrrr rrrr)rSrTrrT)F)rSrTrgrrrT)__name__ __module__ __qualname____doc__r6rfrh __classcell__r#r#rQr$rs w7r) __future__rtypingrrBtorch.nnr9 monai.utilsrr=r`rGrr#r#r#r$s