U Ph @srddlmZddlZddlZddlmZmZddlmZm Z ddl Z ddl m Z ddlmZddlmZddlmZddlmZmZmZmZmZdd lmZmZmZdd lmZm Z dd l!m"Z"m#Z#m$Z$m%Z%m&Z&m'Z'dd l(m)Z)dd l*m+Z+m,Z,ddl-m.Z.m/Z/ddl0m1Z1ee2Z3ddddddZ4ddZ5dddddZ6Gddde Z7e/dd e.d!Gd"d#d#ee7Z8dS)$) annotationsN)MappingMutableMapping)Anycast) DataAnalyzer) get_logger) SegSummarizer)BundleWorkflowConfigComponent ConfigItem ConfigParserConfigWorkflow)SupervisedEvaluatorSupervisedTrainerTrainer) ClientAlgoClientAlgoStats) ExtraItems FiltersTypeFlPhase FlStatistics ModelType WeightType)ExchangeObject)copy_model_stateget_state_dict) min_version require_pkg) DataStatsKeysrrztuple[MutableMapping, int])global_weightslocal_var_dictreturnc Cs|}d}|D]v}||kr||}z,tt|||j}|||<|d7}Wqtk r}ztd|d|W5d}~XYqXq||fS)zAHelper function to convert global weights to local weights formatrzConvert weight from z failed.N)keystorchreshape as_tensorshape Exception ValueError)r r!Z model_keys n_convertedvar_nameweightser/O/home/dell461/cl/sdc2/HISourceFinder-master-l/src/monai/fl/client/monai_algo.pyconvert_global_weights%s &r1cCs|dkrtd|dkr tdi}d}|D]V}||kr:q,||||||<|d7}tt||r,td|dq,|dkrtd|S)Nz>Cannot compute weight differences if `global_weights` is None!z>Cannot compute weight differences if `local_var_dict` is None!rr#z Weights for z became NaN...zNo weight differences computed!)r*cpur%anyisnan RuntimeError)r r!Z weight_diffZn_diffnamer/r/r0compute_weight_diff8s r7r None)parserr"cCs8d|kr4|dD]"}t|rd|dkrd|d<qdS)Nzvalidate#handlersCheckpointLoader_target_T _disabled_)r is_instantiable)r9hr/r/r0disable_ckpt_loadersMs    r?c@sdeZdZdZddddddd d d d Zdd dZddddddZdddZeddZ ddZ dS)MonaiAlgoStatsa7 Implementation of ``ClientAlgoStats`` to allow federated learning with MONAI bundle configurations. Args: bundle_root: directory path of the bundle. config_train_filename: bundle training config path relative to bundle_root. Can be a list of files; defaults to "configs/train.json". only useful when `workflow` is None. config_filters_filename: filter configuration file. Can be a list of files; defaults to `None`. data_stats_transform_list: transforms to apply for the data stats result. histogram_only: whether to only compute histograms. Defaults to False. workflow: the bundle workflow to execute, usually it's training, evaluation or inference. if None, will create an `ConfigWorkflow` internally based on `config_train_filename`. configs/train.jsonNFstrstr | list | None list | NoneboolBundleWorkflow | None) bundle_rootconfig_train_filenameconfig_filters_filenamedata_stats_transform_listhistogram_onlyworkflowcCst|_||_||_||_d|_d|_||_||_d|_|dk rjt |t sPt d| dkrdt d||_d|_ d|_d|_tj|_d|_dS)Ntrainevalz.workflow must be a subclass of BundleWorkflow.z"workflow doesn't specify the type.)loggerrGrHrItrain_data_key eval_data_keyrJrKrL isinstancer r*get_workflow_type client_nameapp_rootpost_statistics_filtersrIDLEphase dataset_root)selfrGrHrIrJrKrLr/r/r0__init__ds(   zMonaiAlgoStats.__init__cCs|dkr i}|tjd|_|tjd}|jd|jd|tjd|_t j |j|j |_ |j dkr||j}t|d|dd|_ |j |j |j _ |j ||j}t}t|dkr|||jtjtdtjd |_|jd |jd dS)  Initialize routine to parse configuration files and extract main components such as trainer, evaluator, and filters. Args: extra: Dict with additional information that should be provided by FL system, i.e., `ExtraItems.CLIENT_NAME`, `ExtraItems.APP_ROOT` and `ExtraItems.LOGGING_FILE`. You can diable the logging logic in the monai bundle by setting {ExtraItems.LOGGING_FILE} to False. Nnoname Initializing  ...rOrM config_file meta_file logging_file workflow_typerdefault Initialized .)getr CLIENT_NAMErU LOGGING_FILErPinfoAPP_ROOTrVospathjoinrGrL_add_config_filesrHr initializerIr len read_configget_parsed_contentrPOST_STATISTICS_FILTERSr rW)r[extrardconfig_train_filesconfig_filter_files filter_parserr/r/r0rss6          zMonaiAlgoStats.initialize dict | Nonerrxr"c Cs|dkrtd|jjrxtj|_|jd|jjtj |krLtdn |tj }tj |krjtdn |tj }i}|j |jj |j ||tj|jdd\}}|r||j |id}d}|jjdk r|j |jj|j||tj|jdd\}}n |jd |r||j|i|rF|rF|||g||} |tj| it|d } |jdk rt|jD]} | | |} qb| Std dS) aX Returns summary statistics about the local data. Args: extra: Dict with additional information that can be provided by the FL system. Both FlStatistics.HIST_BINS and FlStatistics.HIST_RANGE must be provided. Returns: stats: ExchangeObject with summary statistics. Nz`extra` has to be setzComputing statistics on z1FlStatistics.NUM_OF_BINS not specified in `extra`z0FlStatistics.HIST_RANGE not specified in `extra`ztrain_data_stats.yaml)datadata_key hist_bins hist_range output_pathzeval_data_stats.yamlz0the datalist doesn't contain validation section.) statisticszdata_root not set!)r*rL dataset_dirrGET_DATA_STATSrYrPrmr HIST_BINS HIST_RANGE_get_data_key_statstrain_dataset_datarQrorprqrVupdateval_dataset_datarRwarning_compute_total_stats TOTAL_DATArrW) r[rxrrZ stats_dictZtrain_summary_statsZtrain_case_statsZeval_summary_statsZeval_case_statstotal_summary_statsstats_filterr/r/r0get_data_statss^                zMonaiAlgoStats.get_data_statsc Cst||i|jj||||jd}|j|jd|d|j|j|d}|t j }t j |t j t jt|t jt|t|i} | |fS)N)datalistdatarootrrrrKz compute data statistics on z...)transform_listkey)rrLrrKrPrmrUget_all_case_statsrJrBY_CASEr DATA_STATSSUMMARY DATA_COUNTrt FAIL_COUNT) r[r~rrrranalyzerZ all_statsZ case_stats summary_statsr/r/r0rs&  z"MonaiAlgoStats._get_data_key_statscCsRg}|D] }||7}qtdddd||d}||}tj|tjt|tjdi}|S)NimagelabelT)averagedo_ccprrr)r summarizerrrrtr)Zcase_stats_listsrrZtotal_case_statsZcase_stats_list summarizerrrr/r/r0rs(  z#MonaiAlgoStats._compute_total_statscCsg}|rt|tr*|tj|j|nht|trz|D]>}t|tr^|tj|j|q8tdt |d|q8ntdt |d||S)Nz/Expected config file to be of type str but got z: z8Expected config files to be of type str or list but got ) rSrBappendrorprqrGlistr*type)r[ config_filesfilesfiler/r/r0rr$s   z MonaiAlgoStats._add_config_files)rANNFN)N)N)N) __name__ __module__ __qualname____doc__r\rsrr staticmethodrrrr/r/r/r0r@Us (M  r@ignitez0.4.10)pkg_nameversionversion_checkerc@seZdZdZd*d d d dddddd ddddd dddddZd+ddZd,ddddddZd-ddZd.dddddd Zd/d!d"Z d0ddd#d$d%Z d&d'Z d(d)Z dS)1 MonaiAlgoa Implementation of ``ClientAlgo`` to allow federated learning with MONAI bundle configurations. Args: bundle_root: directory path of the bundle. local_epochs: number of local epochs to execute during each round of local training; defaults to 1. send_weight_diff: whether to send weight differences rather than full weights; defaults to `True`. config_train_filename: bundle training config path relative to bundle_root. can be a list of files. defaults to "configs/train.json". only useful when `train_workflow` is None. train_kwargs: other args of the `ConfigWorkflow` of train, except for `config_file`, `meta_file`, `logging_file`, `workflow_type`. only useful when `train_workflow` is None. config_evaluate_filename: bundle evaluation config path relative to bundle_root. can be a list of files. if "default", ["configs/train.json", "configs/evaluate.json"] will be used. this arg is only useful when `eval_workflow` is None. eval_kwargs: other args of the `ConfigWorkflow` of evaluation, except for `config_file`, `meta_file`, `logging_file`, `workflow_type`. only useful when `eval_workflow` is None. config_filters_filename: filter configuration file. Can be a list of files; defaults to `None`. disable_ckpt_loading: do not use any CheckpointLoader if defined in train/evaluate configs; defaults to `True`. best_model_filepath: location of best model checkpoint; defaults "models/model.pt" relative to `bundle_root`. final_model_filepath: location of final model checkpoint; defaults "models/model_final.pt" relative to `bundle_root`. save_dict_key: If a model checkpoint contains several state dicts, the one defined by `save_dict_key` will be returned by `get_weights`; defaults to "model". If all state dicts should be returned, set `save_dict_key` to None. data_stats_transform_list: transforms to apply for the data stats result. eval_workflow_name: the workflow name corresponding to the "config_evaluate_filename", default to "train" as the default "config_evaluate_filename" overrides the train workflow config. this arg is only useful when `eval_workflow` is None. train_workflow: the bundle workflow to execute training, if None, will create a `ConfigWorkflow` internally based on `config_train_filename` and `train_kwargs`. eval_workflow: the bundle workflow to execute evaluation, if None, will create a `ConfigWorkflow` internally based on `config_evaluate_filename`, `eval_kwargs`, `eval_workflow_name`. r#TrANrgmodels/model.ptmodels/model_final.ptmodelrMrBintrErCr|z str | NonerDrF)rG local_epochssend_weight_diffrH train_kwargsconfig_evaluate_filename eval_kwargsrIdisable_ckpt_loadingbest_model_filepathfinal_model_filepath save_dict_keyrJeval_workflow_nametrain_workflow eval_workflowcCsJt|_||_||_||_||_|dkr*in||_|dkr@ddg}||_|dkrRin||_||_| |_ t j | t j | i|_ | |_| |_||_d|_d|_|dk rt|tr|dkrtdtjd||_|dk rt|tr|dkrtd||_d|_d|_d|_d|_d|_d|_d|_d|_d |_ d|_!t"j#|_$d|_%d|_&dS) NrgrAzconfigs/evaluate.jsonrMz6train workflow must be BundleWorkflow and set type in riz3train workflow must be BundleWorkflow and set type.rOr)'rPrGrrrHrrrrIrr BEST_MODEL FINAL_MODELmodel_filepathsrrJrrrrSr rTr*supported_train_type stats_senderrVr{trainer evaluator pre_filterspost_weight_filterspost_evaluate_filtersiter_of_start_timer rrXrYrUrZ)r[rGrrrHrrrrIrrrrrJrrrr/r/r0r\ZsR zMonaiAlgo.__init__cCs*||dkri}|tjd|_|tjd}td}|j d|jd|tj d|_ t j |j |j|_|jdkr|jdk r||j}d|jkr|jd||jd<tf|d|d d |j|_|jdk rX|j|j|j_|j|j_|jr t|jtr t|jjd |j|jj|_t|jtsXtd t|jd |j dkr|j!dk r||j!}d|j"kr|jd||j"d<tf|d||j#d |j"|_ |j dk r8|j |j|j _|jrt|j trt|j jd |j |j j$|_$t|j$t%s8tdt|j$d ||j&}t'|_(t)|dkrf|j(*||tj+|j,|_,|j,dk r|j,-|j|j,-|j$|j(j.t/j0t1dt/j0d|_2|j(j.t/j3t1dt/j3d|_4|j(j.t/j5t1dt/j5d|_6|j(j.t/j7t1dt/j7d|_8|j d|jd dS)r]Nr^z %Y%m%d_%H%M%Sr_r`rOrun_name_rMra)r9z,trainer must be SupervisedTrainer, but got: riz0evaluator must be SupervisedEvaluator, but got: rrfrh)9_set_cuda_devicerjrrkrUrltimestrftimerPrmrnrVrorprqrGrrHrrrrrsr max_epochsrrSr?r9rrr*rrrrrrrrIr r{rtru STATS_SENDERrattachrvr PRE_FILTERSr rPOST_WEIGHT_FILTERSrPOST_EVALUATE_FILTERSrrwrW)r[rxrd timestampryZconfig_eval_filesrzr/r/r0rss                        zMonaiAlgo.initializerr8)r~rxr"cCs4||dkri}t|ts0tdt||jdkrBtd|jdk rb|jD]}|||}qRtj|_ |j d|j dt |jj}ttt|j|d\|_}||j|||jjj|j|jj_|jjj|_ttt|j|jjd\}}}t|dkr|j d |j d |j d |jdS) z Train on client's local data. Args: data: `ExchangeObject` containing the current global model weights. extra: Dict with additional information that can be provided by the FL system. N0expected data to be ExchangeObject but received z self.trainer should not be None.Load weights...r r!srcdstrNo weights loaded!Start z training...) rrSrr*rrrrTRAINrYrPrmrUrnetworkr1rdictr-r _check_convertedstateepochrr iterationrrrrtrrun)r[r~rxrr!r+r updated_keysr/r/r0rMs2           zMonaiAlgo.trainc Cs||dkri}tj|_tj|kr|tj}t|tsNt dt |||j krt j |jtt|j |}t j |st d|tj|dd}t|tr|j|kr||j}tj}i}|jd|d|dnt d |d |j n|jrt|jj}|D]}||||<qtj}|j }|jj!j"|j#|t$j%<|j&r~t'|j(|d }tj)}|jd n |jd nd}d}t}t|tst d|t*|d||d}|j+dk r|j+D]} | ||}q|S)av Returns the current weights of the model. Args: extra: Dict with additional information that can be provided by the FL system. Returns: return_weights: `ExchangeObject` containing current weights (default) or load requested model type from disk (`ModelType.BEST_MODEL` or `ModelType.FINAL_MODEL`). NzEExpected requested model type to be of type `ModelType` but received z#No best model checkpoint exists at r2) map_locationz Returning z checkpoint weights from rizRequested model type z% not specified in `model_filepaths`: rz%Returning current weight differences.zReturning current weights.zstats is not a dict, )r-optim weight_typer),rr GET_WEIGHTSrYr MODEL_TYPErjrSrr*rrrorprqrGrrBisfiler%loadrrrWEIGHTSrPrmrrrr$r2 get_statsrrrrNUM_EXECUTED_ITERATIONSrr7r WEIGHT_DIFFrr) r[rx model_typeZ model_pathr-Z weigh_typerkZreturn_weightsrr/r/r0 get_weights#sd              zMonaiAlgo.get_weightsc Cs`||dkri}t|ts0tdt||jdkrBtd|jdk rb|jD]}|||}qRtj|_ |j d|j dt |jj}ttt|j|d\}}||j||t||jjd\}}}t|dkr|j d |j d |j d t|jtr|j|jjjd n |jt|jjjd } |jdk r\|jD]}|| |} qJ| S)aK Evaluate on client's local data. Args: data: `ExchangeObject` containing the current global model weights. extra: Dict with additional information that can be provided by the FL system. Returns: return_metrics: `ExchangeObject` containing evaluation metrics. Nrz"self.evaluator should not be None.rrrrrrrz evaluating...r#)metrics)rrSrr*rrrrEVALUATErYrPrmrUrrr1rrr-rrrtrrrrrrrr) r[r~rxrr!r r+rrZreturn_metricsr/r/r0evaluaters<              zMonaiAlgo.evaluatecCsz|jd|jd|jdt|jtrJ|jd|jd|jt|jtrv|jd|jd|jdS)z Abort the training or evaluation. Args: extra: Dict with additional information that can be provided by the FL system. z Aborting  during  phase. trainer... evaluator...N) rPrmrUrYrSrrZ interruptrr[rxr/r/r0aborts   zMonaiAlgo.abortr}cCs|jd|jd|jdt|jtrJ|jd|jd|jt|jtrv|jd|jd|j|j dk r|j |j dk r|j dS)z Finalize the training or evaluation. Args: extra: Dict with additional information that can be provided by the FL system. z Terminating rrrrN) rPrmrUrYrSrr terminaterrfinalizerrr/r/r0rs       zMonaiAlgo.finalizecCsB|dkr tdt|n|jd|dt|ddS)Nrz;No global weights converted! Received weight dict keys are z Converted z global variables to match z local variables.)r5rr$rPrmrt)r[r r!r+r/r/r0rszMonaiAlgo._check_convertedcCs*tr&ttjd|_tj|jdS)N LOCAL_RANK) distis_initializedrroenvironrankr%cuda set_device)r[r/r/r0rszMonaiAlgo._set_cuda_device)r#TrANrgNNTrrrNrMNN)N)N)N)N)N)N) rrrrr\rsrMrrrrrrr/r/r/r0r6s2%.A `( O0  r)9 __future__rrorcollections.abcrrtypingrrr%torch.distributed distributedr"monai.apps.auto3dseg.data_analyzerrmonai.apps.utilsrmonai.auto3dsegr monai.bundler r r r r monai.enginesrrrZmonai.fl.clientrrmonai.fl.utils.constantsrrrrrrmonai.fl.utils.exchange_objectrmonai.networks.utilsrr monai.utilsrrmonai.utils.enumsrrrPr1r7r?r@rr/r/r/r0 s2        b