o  i@szddlmZddlmZmZddlZddlmZmZddl m Z m Z ddl m Z ddlmZddlmZGd d d ZdS) ) annotations)CallableSequenceN)decollate_batchlist_data_collate)SupervisedEvaluatorSupervisedTrainer)IterationEvents)Compose) CommonKeysc@s(eZdZdZ ddd dZdddZdS) 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 probability transformsSequence[Callable] | Callablemax_interactionsinttrainboolkey_probabilitystrreturnNonecCs.t|ts t|}||_||_||_||_dS)N) isinstancer rrrr)selfrrrrra/home/dell461/cl/sdc2/last_ska_mid/HISourceFinder-master-l/src/monai/apps/deepgrow/interaction.py__init__*s  zInteraction.__init__engine'SupervisedTrainer | SupervisedEvaluator batchdatadict[str, torch.Tensor]dictc CsN|durtdt|jD]}||\}}||jj}|tj |j t /|jrMt d|||j }Wdn1sGwYn|||j }Wdn1s^wY|tj|tj|it|dd}tt|D]}|jrdd|j|nd|||j<|||||<q}t|}q |||S)Nz.Must provide batch data for current iteration.cudaT)detachg?) ValueErrorranger prepare_batchtostatedevice fire_eventr INNER_ITERATION_STARTEDnetworkevaltorchno_gradampautocastinfererINNER_ITERATION_COMPLETEDupdater PREDrlenrrrr _iteration) rrrjinputs_ predictionsbatchdata_listirrr__call__9s2         zInteraction.__call__N)r ) rrrrrrrrrr)rrrr rr!)__name__ __module__ __qualname____doc__rr>rrrrr s  r ) __future__rcollections.abcrrr. monai.datarr monai.enginesrrmonai.engines.utilsr monai.transformsr monai.utils.enumsr r rrrrs