U {Ph>@sddlmZddlZddlZddlZddlmZddlmZm Z ddl m Z ddl m Z mZmZede jed\ZZerdd lmZdd lmZmZn._DiskSaverzK Enhance the DiskSaver to support fixed filename. Nrrdirnamefilenamecstj|ddd||_dS)NF)r)Z require_emptyatomic)super__init__r*)selfr)r* __class__T/home/dell461/cl/sdc2/HISourceFinder-master-l/src/monai/handlers/checkpoint_saver.pyr-sz5CheckpointSaver.__init__.._DiskSaver.__init__rzMapping | Noner) checkpointr*metadatar'cs&|jdk r|j}tj|||ddS)N)r3r*r4)r*r,__call__)r.r3r*r4r/r1r2r5s z5CheckpointSaver.__init__.._DiskSaver.__call__)r*r'cs"|jdk r|j}tj|ddS)N)r*)r*r,remove)r.r*r/r1r2r6s z3CheckpointSaver.__init__.._DiskSaver.remove)N)N)__name__ __module__ __qualname____doc__r-r5r6 __classcell__r1r1r/r2 _DiskSaver|sr<r renginer'cSs|jjSN)state iterationr>r1r1r2 _final_funcsz-CheckpointSaver.__init__.._final_funcr(Zfinal_iteration)to_save save_handlerfilename_prefixscore_function score_namecsvttr}n&t|jdr&|jj}ntdd|jj|}t|sft d|d|ddSrndnd|S)Nrz>Incompatible values: save_key_metric=True and key_metric_name=.zkey metric is not a scalar value, skip metric comparison and don't save a model.please use other metrics as key metric, or change the `reduction` mode to 'mean'.got metric: =r) isinstancerhasattrr@r ValueErrormetricsrwarningswarn)r> metric_namemetric)rr#r1r2 _score_funcs     z-CheckpointSaver.__init__.._score_funcrzSif using fixed filename to save the best metric model, we should only save 1 model. key_metric)rDrErFrGrHr& include_selfZgreater_or_equalcsjr|jjS|jjSr?)r$r@epochrArB)r.r1r2_interval_funcsz0CheckpointSaver.__init__.._interval_func)r)rWrA)rDrErFrGrHr&)AssertionErrorrlenrlogging getLoggerloggerr$r%_final_checkpoint_key_metric_checkpoint_interval_checkpoint_name_final_filenamer r rN)r.rrrrrrrrrr r!r"r#r$r%r&r<rCrTrXr1)rr#r.r2r-Zs`    zCheckpointSaver.__init__) state_dictr'cCs&|jdk r|j|n tddS)a Utility to resume the internal state of key metric tracking list if configured to save checkpoints based on the key metric value. Note to set `key_metric_save_state=True` when saving the previous checkpoint. Example:: CheckpointSaver( ... save_key_metric=True, key_metric_save_state=True, # config to also save the state of this saver ).attach(engine) engine.run(...) # resumed training with a new CheckpointSaver saver = CheckpointSaver(save_key_metric=True, ...) # load the previous key metric tracking list into saver CheckpointLoader("/test/model.pt"), {"checkpointer": saver}).attach(engine) NzFno key metric checkpoint saver to resume the key metric tracking list.)r_load_state_dictrPrQ)r.rcr1r1r2rds zCheckpointSaver.load_state_dictr r=cCs|jdkr|j|_|jdk r<|tj|j|tj|j|j dk rV|tj |j |j dk r|j r|tj |jd|jn|tj|jd|jdS)zg Args: engine: Ignite Engine, it can be a trainer, validator or evaluator. N)every)rar]r^add_event_handlerr Z COMPLETED completedZEXCEPTION_RAISEDexception_raisedr_EPOCH_COMPLETEDmetrics_completedr`r$r%interval_completedITERATION_COMPLETEDr.r>r1r1r2attachs    zCheckpointSaver.attachcCsP|jdk rL|jj}t|dkrL|d}|jj|j|jd|jdS)Nrz)Deleted previous saved final checkpoint: ) r^Z_savedrZpoprEr6r*r]info)r.saveditemr1r1r2_delete_previous_final_ckpts    z+CheckpointSaver._delete_previous_final_ckptcCst|jstd||||jdkr2tt|jdsFtd|jdk rdtj |j |j}n|jj }|j d|dS)zCallback for train or validation/evaluation completed Event. Save final checkpoint if configure save_final is True. Args: engine: Ignite Engine, it can be a trainer, validator or evaluator. 0Error: _final_checkpoint function not specified.Nrp.Error, provided logger has not info attribute.z)Train completed, saved final checkpoint: callabler^rYrsr]rMrbospathjoinrZlast_checkpointrp)r.r>_final_checkpoint_pathr1r1r2rgs     zCheckpointSaver.completed Exception)r>er'cCst|jstd||||jdkr2tt|jdsFtd|jdk rdtj |j |j}n|jj }|j d||dS)aCallback for train or validation/evaluation exception raised Event. Save current data as final checkpoint if configure save_final is True. This callback may be skipped because the logic with Ignite can only trigger the first attached handler for `EXCEPTION_RAISED` event. Args: engine: Ignite Engine, it can be a trainer, validator or evaluator. e: the exception caught in Ignite during engine.run(). rtNrpruz-Exception raised, saved the last checkpoint: rv)r.r>r}r{r1r1r2rhs     z CheckpointSaver.exception_raisedcCs t|jstd||dS)zCallback to compare metrics and save models in train or validation when epoch completed. Args: engine: Ignite Engine, it can be a trainer, validator or evaluator. z5Error: _key_metric_checkpoint function not specified.N)rwr_rYrmr1r1r2rj3s z!CheckpointSaver.metrics_completedcCsvt|jstd|||jdkr*tt|jds>td|jr\|jd|jjn|jd|jj dS)zCallback for train epoch/iteration completed Event. Save checkpoint if configure save_interval = N Args: engine: Ignite Engine, it can be a trainer, validator or evaluator. z3Error: _interval_checkpoint function not specified.NrpruzSaved checkpoint at epoch: zSaved checkpoint at iteration: ) rwr`rYr]rMr$rpr@rWrArmr1r1r2rk=s    z"CheckpointSaver.interval_completed)NrFNFNrNFFFTrN) r7r8r9r:r-rdrnrsrgrhrjrkr1r1r1r2r"s.;0v r) __future__rr[rxrPcollections.abcrtypingrr monai.configr monai.utilsrrr OPT_IMPORT_VERSIONr _ ignite.enginer Zignite.handlersr r rr1r1r1r2 s