U Ph@szddlmZddlmZmZddlZddlmZmZddl m Z m Z ddl m Z ddlmZddlmZGd d d ZdS) ) annotations)CallableSequenceN)decollate_batchlist_data_collate)SupervisedEvaluatorSupervisedTrainer)IterationEvents)Compose) CommonKeysc@s:eZdZdZdddddddd d Zd d d dddZdS) Interactiona Ignite process_function used to introduce interactions (simulation of clicks) for Deepgrow Training/Evaluation. For more details please refer to: https://pytorch.org/ignite/generated/ignite.engine.engine.Engine.html. This implementation is based on: Sakinis et al., Interactive segmentation of medical images through fully convolutional neural networks. (2019) https://arxiv.org/abs/1903.08205 Args: transforms: execute additional transformation during every iteration (before train). Typically, several Tensor based transforms composed by `Compose`. max_interactions: maximum number of interactions per iteration train: training or evaluation key_probability: field name to fill probability for every interaction probabilityzSequence[Callable] | CallableintboolstrNone) transformsmax_interactionstrainkey_probabilityreturncCs.t|tst|}||_||_||_||_dS)N) isinstancer rrrr)selfrrrrrT/home/dell461/cl/sdc2/HISourceFinder-master-l/src/monai/apps/deepgrow/interaction.py__init__*s  zInteraction.__init__z'SupervisedTrainer | SupervisedEvaluatorzdict[str, torch.Tensor]dict)engine batchdatarc Cs*|dkrtdt|jD]}||\}}||jj}|tj |j t B|jrt jj|||j }W5QRXn|||j }W5QRX|tj|tj|it|dd}tt|D]<}|jrdd|j|nd|||j<|||||<qt|}q|||S)Nz.Must provide batch data for current iteration.T)detachg?) ValueErrorranger prepare_batchtostatedevice fire_eventr INNER_ITERATION_STARTEDnetworkevaltorchno_gradampcudaautocastinfererINNER_ITERATION_COMPLETEDupdater PREDrlenrrrr _iteration) rrrjinputs_ predictionsbatchdata_listirrr__call__9s*       zInteraction.__call__N)r )__name__ __module__ __qualname____doc__rr;rrrrr sr ) __future__rcollections.abcrrr* monai.datarr monai.enginesrrmonai.engines.utilsr monai.transformsr monai.utils.enumsr r rrrr s