U Phn=@s(ddlmZddlZddlmZddlZddlmZddlmZddl m Z e dZ e ddd dZ e d d d dZ e d d d dZe d dd dZdddddddgZGdddejZGdddejZGdddejZGdddejZGdddejZGdddeZGdddejjZdS)) annotationsN)Sequence)nn)PathLike)optional_import transformersload_tf_weights_in_bert)nameztransformers.utils cached_filez&transformers.models.bert.modeling_bertBertEmbeddings BertLayerBertPreTrainedModel BertAttention BertOutputBertMixedLayerPooler MultiModal Transchexcs<eZdZdZddfdd ZddZedd d ZZS)r zModule to load BERT pre-trained weights. Based on: LXMERT https://github.com/airsplay/lxmert BERT (pytorch-transformer) https://github.com/huggingface/transformers NonereturncstdSN)super__init__)selfinputskwargs __class__R/home/dell461/cl/sdc2/HISourceFinder-master-l/src/monai/networks/nets/transchex.pyr)szBertPreTrainedModel.__init__cCsxt|tjtjfr*|jjjd|jjdn(t|t jj rR|j j |jj dt|tjrt|j dk rt|j j dS)N)meanstd?) isinstancerLinear Embeddingweightdatanormal_configinitializer_rangetorch LayerNormbiaszero_fill_)rmodulerrr init_bert_weights,s z%BertPreTrainedModel.init_bert_weightsNFbert-base-uncasedpytorch_model.binc sZt|| |d} |||||f| | } dkrL|sLtj| tjsDdndd|rZt| | Sg}g}D]H}d}d|kr|dd}d|kr|dd}|rj||||qjt ||D]\}} ||<qgggt dd dk r_ dfd d d }t| d sJtd dDrJd}| |d| S)N) cache_dircpu) map_locationgammar(betar/ _metadatac shdkr in|ddi}|||d|jD]"\}}|dk r@|||dq@dS)NT.)get_load_from_state_dict_modulesitems)r2prefixlocal_metadatar child error_msgsloadmetadata missing_keys state_dictunexpected_keysrr rH_s z1BertPreTrainedModel.from_pretrained..loadbertcss|]}|dVqdS)bert.N) startswith).0srrr isz6BertPreTrainedModel.from_pretrained..rN)rC)r<)r r-rHcuda is_availablerkeysreplaceappendzippopgetattrcopyr;hasattrany)clsnum_language_layersnum_vision_layersnum_mixed_layers bert_configrKr6Zfrom_tfpath_or_repo_idfilenamerrZ weights_pathmodelZold_keysnew_keyskeynew_keyold_keyZ start_prefixrrFr from_pretrained5s@          $ z#BertPreTrainedModel.from_pretrained)NNFr4r5) __name__ __module__ __qualname____doc__rr3 classmethodrj __classcell__rrrr r s cs6eZdZdZddfdd ZddZdd ZZS) rzsBERT attention layer. Based on: BERT (pytorch-transformer) https://github.com/huggingface/transformers rrcszt|j|_t|j|j|_|j|j|_t|j|j|_ t|j|j|_ t|j|j|_ t |j |_dSr)rrnum_attention_headsint hidden_sizeattention_head_size all_head_sizerr&queryrgvalueDropoutattention_probs_dropout_probdropoutrr+rrr rus zBertAttention.__init__cCs6|dd|j|jf}|j|}|ddddS)Nr=r)sizerqrtviewpermute)rxZ new_x_shaperrr transpose_for_scoress z"BertAttention.transpose_for_scoresc Cs||}||}||}||}||}||}t||dd} | t|j } | t j dd| } t| |} | dddd} | dd|jf} | j| } | S)Nr=)dimrr|r}r~)rvrgrwrr-matmul transposemathsqrtrtrzrSoftmaxr contiguousrrur) r hidden_statescontextZmixed_query_layerZmixed_key_layerZmixed_value_layer query_layerZ key_layerZ value_layerZattention_scoresattention_probsZ context_layerZnew_context_layer_shaperrr forwards        zBertAttention.forward)rkrlrmrnrrrrprrrr ros cs.eZdZdZddfdd ZddZZS)rzpBERT output layer. Based on: BERT (pytorch-transformer) https://github.com/huggingface/transformers rrcsBtt|j|j|_tjj|jdd|_t|j |_ dS)N-q=)eps) rrrr&rsdenser-r.rxhidden_dropout_probrzr{rrr rs zBertOutput.__init__cCs&||}||}|||}|Sr)rrzr.)rr input_tensorrrr rs  zBertOutput.forwardrkrlrmrnrrrprrrr rscs.eZdZdZddfdd ZddZZS)rzyBERT cross attention layer. Based on: BERT (pytorch-transformer) https://github.com/huggingface/transformers rrcs6tt||_t||_t||_t||_dSr)rrratt_xroutput_xatt_youtput_yr{rrr rs     zBertMixedLayer.__init__cCs0|||}|||}||||||fSr)rrrr)rryrrrrr rs  zBertMixedLayer.forwardrrrrr rscs.eZdZdZddfdd ZddZZS)rzpBERT pooler layer. Based on: BERT (pytorch-transformer) https://github.com/huggingface/transformers rrcs&tt|||_t|_dSr)rrrr&rTanh activation)rrsrrr rs zPooler.__init__cCs(|dddf}||}||}|SNr)rr)rrZfirst_token_tensorZ pooled_outputrrr rs  zPooler.forwardrrrrr rscs8eZdZdZddddddfdd Zd d d ZZS) rz? Multimodal Transformers From Pretrained BERT Weights" rrdictr)r_r`rarbrcsttdtf|_tj_tfddt |D_ tfddt |D_ tfddt |D_ jdS)z Args: num_language_layers: number of language transformer layers. num_vision_layers: number of vision transformer layers. bert_config: configuration for bert language transformer encoder. objcsg|]}tjqSrr r+rP_rrr sz'MultiModal.__init__..csg|]}tjqSrrrrrr rscsg|]}tjqSr)rr+rrrr rsN)rrtypeobjectr+r embeddingsr ModuleListrangelanguage_encodervision_encoder mixed_encoderapplyr3)rr_r`rarbrrr rs  zMultiModal.__init__NcCsb|||}|jD]}||dd}q|jD]}|||d}q,|jD]}|||\}}qF||fSr)rrrr)r input_idstoken_type_ids vision_featsattention_maskZlanguage_featureslayerrrr rs    zMultiModal.forward)NNNrrrrr rsc"speZdZdZd#ddddddddddddddddddddddddddddddddd fdd Zd$d!d"ZZS)%rz TransChex based on: "Hatamizadeh et al.,TransCheX: Self-Supervised Pretraining of Vision-Language Transformers for Chest X-ray Analysis" r!皙?Fgelu{Gz? rrM rabsolute4.10.2r|T:wr4r5rrzSequence[int] | intzint | tuple[int, int]floatboolstrzstr | PathLiker) in_channelsimg_size patch_size num_classesr_r`rarsdrop_outrygradient_checkpointing hidden_actrr,intermediate_sizelayer_norm_epsmax_position_embeddings model_typerqnum_hidden_layers pad_token_idposition_embedding_typetransformers_versiontype_vocab_size use_cache vocab_sizechunk_size_feed_forward is_decoderadd_cross_attentionrcrdrc !s8t| d| | | |||||||||||||||||d} d| krPdksZntd|d|ddks|d|ddkrtdtj|||| ||d|_||_|d|jd|d|jd|_tj |||j|jd|_ t ||_ t td|j||_t|d |_tj| |_tj|||_dS) a Args: in_channels: dimension of input channels. img_size: dimension of input image. patch_size: dimension of patch size. num_classes: number of classes if classification is used. num_language_layers: number of language transformer layers. num_vision_layers: number of vision transformer layers. num_mixed_layers: number of mixed transformer layers. drop_out: fraction of the input units to drop. path_or_repo_id: This can be either: - a string, the *model id* of a model repo on huggingface.co. - a path to a *directory* potentially containing the file. filename: The name of the file to locate in `path_or_repo`. The other parameters are part of the `bert_config` to `MultiModal.from_pretrained`. Examples: .. code-block:: python # for 3-channel with image size of (224,224), patch size of (32,32), 3 classes, 2 language layers, # 2 vision layers, 2 mixed modality layers and dropout of 0.2 in the classification head net = Transchex(in_channels=3, img_size=(224, 224), num_classes=3, num_language_layers=2, num_vision_layers=2, num_mixed_layers=2, drop_out=0.2) N)ryZclassifier_dropoutrrrrsr,rrrrrqrrrrrrrrrrrr}z'dropout_rate should be between 0 and 1.z+img_size should be divisible by patch_size.)r_r`rarbrcrd)r out_channels kernel_sizestride)rs)rr ValueErrorrrj multimodalr num_patchesrConv2d vision_projr.norm_vision_pos Parameterr-zeros pos_embed_visrpoolerrxdropr&cls_head)!rrrrrr_r`rarsrryrrrr,rrrrrqrrrrrrrrrrrcrdrbrrr rsbB ( &  zTranschex.__init__Nc Cst|dd}|jt|jd}d|d}||d dd}| |}||j }|j ||||d\}}| |}|||}|S)Nr}r|)dtyper$g)rrrr)r- ones_like unsqueezetonext parametersrrflattenrrrrrrr) rrrrrZhidden_state_langZhidden_state_visZpooled_featureslogitsrrr rjs     zTranschex.forward)rr!rFrrrrrrrMrrrrrr|TrrFFr4r5)NNrrrrr rs6Ru) __future__rrcollections.abcrr-rmonai.config.type_definitionsr monai.utilsrrrr r r __all__Moduler rrrrrrrrrr  s&     O&"