o & ik@sddlmZddlZddlZddlZddlZddlmZddlm Z ddl m Z ddl m Z mZmZmZmZddlZddlmZddlmZmZmZdd lmZed \ZZerldd lmZdd lm Z m!Z!m"Z"m#Z#m$Z$ed \Z%Z&edd\Z'Z(ed\Z)Z(e*Z+ddZ,ddZ-ddZ.Gddde/Z0GdddZ1ddZ2ddZ3d7d%d&Z4Gd'd(d(Z5d)d*Z6   d8d9d5d6Z7dS):) annotationsN) OrderedDict)Path) MethodType)AnyDictListTupleUnion) get_logger)add_casts_around_normsconvert_to_onnxget_profile_shapes)optional_import polygraphy)bytes_from_path) CreateConfigProfileengine_bytes_from_networkengine_from_bytesnetwork_from_onnx_pathZtensorrttorch_tensorrtz1.4.0z cuda.cudartcCs<tjtjtjtjtjtjtjtjtjtjtjtjtjtjiSN) trtint32torchfloat32float16bfloat16int64int8boolr"r"]/home/dell461/cl/sdc2/last_ska_mid/HISourceFinder-master-l/src/monai/networks/trt_compiler.pytrt_to_torch_dtype_dict1sr$cCs|i}|s|S|D]3}|D].}g}||}tt|dD]}|d||d|kr/||qt|dkr:|||<q q|S)z This method calculates dynamic_axes to use in onnx.export(). Args: profiles: [[min,opt,max],...] list of profile dimensions r)rangelenappend)profiles dynamic_axesprofilekeyaxesvalsir"r"r#get_dynamic_axes=s   r0cCs6|d}|dkrtd|t|dkr|dSdS)z[ Error reporting method for CUDA calls. Args: cuda_ret: CUDA return code. rz CUDA ERROR: N) RuntimeErrorr')Zcuda_reterrr"r"r#cuassertRs  r4c@seZdZdZdS) ShapeErrorzM Exception class to report errors from setting TRT plan input shapes N)__name__ __module__ __qualname____doc__r"r"r"r#r5`sr5c@s4eZdZdZd ddZddZddZd d d ZdS) TRTEnginezK An auxiliary class to implement running of TRT optimized engines NcCs||_|ptd|_|jd|jtt|j|_t|_d|_ |j |_ g|_ g|_ g|_d|_i|_t}t|jjD]6}|j|}|j|tjjkrY|j |qA|j|tjjkrw|j |||j|}|j|qA|jd|jd|j d|j dS)z Loads serialized engine, creates execution context and activates it Args: plan_path: path to serialized TRT engine. logger: optional logger object monai.networks.trt_compilerzLoading TensorRT engine: NrzLoaded TensorRT engine: z . Inputs: z Outputs: ) plan_pathr loggerinforrenginertensorscuda_graph_instanceZcreate_execution_contextcontext input_names output_namesdtypes cur_profile input_tabler$r&Znum_io_tensorsZget_tensor_moderZ TensorIOModeINPUTr(ZOUTPUTZget_tensor_dtype)selfr<r= dtype_dictidxbindingdtyper"r"r#__init__ns2    zTRTEngine.__init__cCs~|j}t|jD]4\}}t||}||jvs"t|j|j|krwq}t|dksLJdS)z Sets input bindings for TRT engine according to feed_dict Args: feed_dict: a dictionary [str->Tensor] stream: CUDA stream to use csTjD]$}j|d}|dur'|}|j}||||qdSr)rCgetrGrTrRZset_input_shaperUrV)rLrXrRrW feed_dictrIr"r#try_set_inputss  z,TRTEngine.set_inputs..try_set_inputsTr1rN) r?rBrFr5Znum_optimization_profilesZset_optimization_profile_async ExceptionZ infer_shapesr')rIr\streameZ last_profiler]Z next_profileleftr"r[r# set_inputss(    zTRTEngine.set_inputsFcCs|rO|jdurtt|j|tt||jS|j|}|s&tdtt|tj j |j|tt |}tt |d|_|j d|jS|j|}tt||sbtd|jS)z Runs TRT engine. Args: stream: CUDA stream to run on use_cuda_graph: use CUDA graph. Note: requires all inputs to be the same GPU memory between calls. NzERROR: inference failed.rzCUDA Graph captured!)rAr4cudartZcudaGraphLaunchZcudaStreamSynchronizerBZexecute_async_v3 ValueErrorZcudaStreamBeginCaptureZcudaStreamCaptureModeZ cudaStreamCaptureModeThreadLocalZcudaStreamEndCaptureZcudaGraphInstantiater=r>r@)rIr_use_cuda_graphZnoerrorgraphr"r"r#infers*     zTRTEngine.inferr)F)r6r7r8r9rNrYrbrgr"r"r"r#r:hs   $r:cCst|tjr|St|Sr) isinstancerTensortensorcuda)dr"r"r# make_tensorsrmcCspi}|D]1}||}|dur5t|tst|tr/tt|D]}t||||d|<qqt|||<q|S)N_)rhrQtupler&r'rm)rC input_exampleunrolled_inputnamevalr/r"r"r# unroll_inputs rtretList[torch.Tensor] output_listsList[List[int]]return3Tuple[Union[torch.Tensor, List[torch.Tensor]], ...]c Cst}d}tt|D]}||}t|dkst|dksJt|dks+|ddkr9g|||R}|d}q |ddkrUg|||||dR}||d}q |ddkrt}t|}tt|d|dD]M}||} t| dkst| dksJt| dks| ddkr|d}g|||R}ql| ddkr|| d}g||||| dR}qltdg|||||dddR}|Sq |S)a) Implements parsing of 'output_lists' arg of trt_compile(). Args: ret: plain list of Tensors output_lists: list of output group sizes: to form some Lists/Tuples out of 'ret' List, this will be a list of group dimensions, like [[], [5], [-1]] for returning Tensor, list of 5 items and dynamic list. Format: [[group_n] | [], ...] [] or group_n == 0 : next output from ret is a scalar group_n > 0 : next output from ret is a list of group_n length group_n == -1: next output is a dynamic list. This entry can be at any position in output_lists, but can appear only once. Returns: Tuple of Union[torch.Tensor, List[torch.Tensor]], according to the grouping in output_lists rr1zTwo -1 lists in outputN)ror&r'rd) rurwgroupscurlglZ rev_groupsZrcurrlZrglr"r"r# parse_groupss:      $rc@s^eZdZdZ              dddZdd Zd d Zd d ZddZddZ dS) TrtCompilerz This class implements: - TRT lazy persistent export - Running TRT with optional fallback to Torch (for TRT engines with limited profiles) fp16onnxNFcCsddg}||vrtd|d|dgd}||vr&td|d|d||_||_||_|du|_|p7g|_|pTRT). This is the most stable and efficient option. 'torch_trt' may not work for some nets. Also AMP must be turned off for it to work. input_names: Optional list of input names. If None, will be read from the function signature. output_names: Optional list of output names. Note: If not None, patched forward() will return a dictionary. output_lists: Optional list of output group sizes: when forward() returns Lists/Tuples, this will be a list of their dimensions, like [[], [5], [-1]] for returning Tensor, list of 5 items and dynamic list. export_args: Optional args to pass to export method. See onnx.export() and Torch-TensorRT docs for details. build_args: Optional args to pass to TRT builder. See polygraphy.Config for details. input_profiles: Optional list of profiles for TRT builder and ONNX export. Each profile is a map of the form : {"input id" : [min_shape, opt_shape, max_shape], ...}. dynamic_batchsize: A sequence with three elements to define the batch size range of the input for the model to be converted. Should be a sequence like [MIN_BATCH, OPT_BATCH, MAX_BATCH]. [note]: If neither input_profiles nor dynamic_batchsize specified, static shapes will be used to build TRT engine. use_cuda_graph: Use CUDA Graph for inference. Note: all inputs have to be the same GPU memory between calls! timestamp: Optional timestamp to rebuild TRT engine (e.g. if config file changes). fallback: Allow to fall back to Pytorch when TRT inference fails (e.g, shapes exceed max profile). r torch_trtz)trt_compile(): 'method' should be one of z, got: .)fp32tf32rbf16z,trt_compile(): 'precision' should be one of NFr;r1)!rdr< precisionmethod return_dictrDrwr)dynamic_batchsize export_args build_argsr?refallbackdisabledr r=inspectgetfullargspecforwardargspecargsdefaultsr&r'rmrC old_forwardospathexistsgetmtimeremove)rImodelr<rrrCrDrwrrZinput_profilesrre timestamprZforward_overrider=Z method_valsZprecision_valsr/rlr"r"r#rN.sH.       ( zTrtCompiler.__init__cCs,i}t|D] \}}|j|}|||<q|Sr)rPrC)rIrpZ trt_inputsr/inp input_namer"r"r#_inputs_to_dicts   zTrtCompiler._inputs_to_dictc Csz:t|j|j|_i}|jjD]}|dr"||jvr"|dd}n|}|||<q||j_|jd|jjWdStyV}z|jd|WYd}~dSd}~ww)zO Loads TRT plan from disk and activates its execution context. __r%NzEngine loaded, inputs:z$Exception while loading the engine: ) r:r<r=r?rC startswithrGr>r^)rIrGrr orig_namer`r"r"r# _load_engines   zTrtCompiler._load_enginec CsV|j}||t|dkr||||jdur|js|j}|j|_z4||jdurX| }t | ||Wdn1sHwY||jdusXJWn$t y}}z|jrq|jd|d|_n|WYd}~nd}~ww|js|js|D]}~qt j||_zj|jdurtYt j} t jj| d} |jt|j|| j|jj| d| t j|jj| j|jd} |j st!| "} |j#rt$| |j#} n t| dkr| d} | WdWS1swYWn$t y"}z|jr|jd|d n|WYd}~nd}~ww|j|i|S) af Main forward method: Builds TRT engine if not available yet. Tries to run TRT engine If exception thrown and self.callback==True: falls back to original Pytorch Args: Passing through whatever args wrapped module's forward() has Returns: Passing through wrapped module's forward() return value(s) rNzFailed to build engine: T)rO)rer1z Exception: z Falling back to Pytorch ...)%rupdater'rr?rrrrcopyrno_grad_build_and_saver^rr=r> parametersrk empty_cachelock_smcurrent_deviceStreamrbrtrC cuda_streamrY wait_streamcurrent_streamrgrerrQvaluesrwr) rIrargvkwargsr new_forwardrr`paramrOr_rur"r"r#rsp            " zTrtCompiler.forwardc Csg}|jD]"}t}|D]\}}|j||d|d|ddq||q|j}|jdk|d<|jdkr>d|d<n |jd krGd|d <|j d |d |j t |t j jgd }t|tdd |i|dS)z[ Builds TRT engine from ONNX file at onnx_path and saves to self.plan_path rr1r%)minoptmaxrrrTrzBuilding TensorRT engine for z: )flagsr))configNr")r)ritemsaddr(rrrr=r>r<rrZOnnxParserFlagZNATIVE_INSTANCENORMrr) rI onnx_pathr)r+pidrsrnetworkr"r"r# _onnx_to_trts       zTrtCompiler._onnx_to_trtc s`jdurdSj}d}t|jdkrRtjg}jdkr%|tjn jdkr0|tj t | }ddfdd|D}t j |d f||d |}nΈjrtjd krbtd td krltdi|D]6\}} fdd} t| t st| trtt| D]} | |d| | | qqrt| tjr| || qrg_tj_tjd kr|djitQ} tj|} tt| d}j !d|dt | "ddj#djd|t$||f|t | "j#d|j !d%|}Wdn 1swY|r.t&j'd(|dSdS)z If TRT engine is not ready, exports model to ONNX, builds TRT engine and saves serialized TRT engine to the disk. Args: input_example: passed to onnx.export() NrrrcSs t||\}}}tj|||dS)N)Z min_shapeZ opt_shapeZ max_shape)rrInput) input_shaperZmin_input_shapeZopt_input_shapeZmax_input_shaper"r"r#get_torch_trt_inputsz8TrtCompiler._build_and_save..get_torch_trt_inputcsg|] }|jjqSr")rRr).0r/)rrIr"r# sz/TrtCompiler._build_and_save..r)Z arg_inputsenabled_precisionsrzEERROR: Both dynamic_batchsize and input_profiles set for TrtCompiler!z&dynamic_batchsize has to have len ==3 csR|j}t|dkr'|dd}dg|dg|dg|g|<dSdS)Nrr1r%)rRr')rrssh)dbsr+r"r# add_profile)s   0z0TrtCompiler._build_and_save..add_profilernr*z model.onnxz Exporting to z: unrolled_inputs= z output_names=z input_names=z export args: )filenamerCrDzExport to ONNX successful.wb))r?rr rrrrr(rrrQrrZconvert_method_to_trt_enginerr'r)rdrrhror&rir0r*rtempfileTemporaryDirectoryrtrCstrrr=r>keysrDr ropenr<write)rIrrprZ engine_bytesrinputsZ tt_inputsrrsrr/tmpdirrqrr")rrr+rIr#rs               zTrtCompiler._build_and_save)rrNNNNNNNFNFNN) r6r7r8r9rNrrrrrr"r"r"r#r&s,  YF rcOs|j|||S)zk Patch function to replace original model's forward() with. Redirects to TrtCompiler.forward() ) _trt_compilerr)rIrrr"r"r# trt_forwardQsrrtorch.nn.Module base_pathrrDict[str, Any] | None submoduleUnion[str, List[str]] | Noner= Any | Nonec sdddddd}|pi|trttrttjrttj|r:t tj |}dvr6t t d|}|d<fdd }fd d |d urmt |t rS|g}|D]}||\} } |t| | |d |qU|S||||Spytdd|S)a Instruments model or submodule(s) with TrtCompiler and replaces its forward() with TRT hook. Note: TRT 10.3 is recommended for best performance. Some nets may even fail to work with TRT 8.x. NVIDIA Volta support (GPUs with compute capability 7.0) has been removed starting with TensorRT 10.5. Review the TensorRT Support Matrix for which GPUs are supported. Args: model: module to patch with TrtCompiler object. base_path: TRT plan(s) saved to f"{base_path}[.{submodule}].plan" path. dirname(base_path) must exist, base_path does not have to. If base_path does point to existing file (e.g. associated checkpoint), that file becomes a dependency - its mtime is added to args["timestamp"]. args: Optional dict : unpacked and passed to TrtCompiler() - see TrtCompiler above for details. submodule: Optional hierarchical id(s) of submodule to patch, e.g. ['image_decoder.decoder'] If None, TrtCompiler patch is applied to the whole model. Otherwise, submodule (or list of) is being patched. logger: Optional logger for diagnostics. Returns: Always returns same model passed in as argument. This is for ease of use in configs. rrZobey)Zbuilder_optimization_levelZprecision_constraints)rrrrcsFt|ds!|j|_t||dfdi}||_tt||_dSdS)Nrz.planr=)hasattrr orig_forwardrrrr)rrwrapper)rr=r"r#wraps ztrt_compile..wrapcsJ|d}|dkr!|d|}t||}||dd}||S||fS)Nrr{r1)findgetattr)parentrrK parent_name)find_subr"r#rs    ztrt_compile..find_subNrr;zSTensorRT and/or polygraphy packages are not available! trt_compile() has no effect.)r trt_importedpolygraphy_importedrrk is_availablerrrintrrrhrrr warning) rrrrr=Z default_argsrrsrsubr")rrr=r# trt_compileYs4     r)rurvrwrxryrz)NNN) rrrrrrrrr=rryr)8 __future__rrrr threading collectionsrpathlibrtypesrtypingrrrr r rmonai.apps.utilsr monai.networks.utilsr r rmonai.utils.modulerrrZpolygraphy.backend.commonrZpolygraphy.backend.trtrrrrrrrrrnrcLockrr$r0r4r^r5r:rmrtrrrrr"r"r"r#sJ           z 2-