o - i=@sddlmZddlZddlmZddlZddlmZddlmZddl m Z e dZ e ddd dZ e d d d dZ e d d d dZe d dd dZgdZGdddejZGdddejZGdddejZGdddejZGdddejZGdddeZGdddejjZdS)) annotationsN)Sequence)nn)PathLike)optional_import transformersload_tf_weights_in_bert)nameztransformers.utils cached_filez&transformers.models.bert.modeling_bertBertEmbeddings BertLayer)BertPreTrainedModel BertAttention BertOutputBertMixedLayerPooler MultiModal TranschexcsBeZdZdZdfdd ZddZe   dd 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 returnNonecstdSN)super__init__)selfinputskwargs __class___/home/dell461/cl/sdc2/last_ska_mid/HISourceFinder-master-l/src/monai/networks/nets/transchex.pyr)szBertPreTrainedModel.__init__cCst|tjtjfr|jjjd|jjdnt|t jj r)|j j |jj dt|tjr<|j dur>|j j dSdSdS)N)meanstd?) isinstancerLinear Embeddingweightdatanormal_configinitializer_rangetorch LayerNormbiaszero_fill_)rmodulerrrinit_bert_weights,s z%BertPreTrainedModel.init_bert_weightsNFbert-base-uncasedpytorch_model.binc s`t|| |d} |||||g| Ri| } dur,|s,tjs"dnd}tj| |dd|r3t| | Sg}g}D]$}d}d|vrI|dd}d|vrS|dd}|r_||||q;t ||D] \}} ||<qegggt d d dur_ dfd d d }t| d stddDrd}| |d| S)N) cache_dircpuT) map_location weights_onlygammar'betar. _metadatac shdurin |ddi}|||d|jD]\}}|dur1|||dq dS)NT.)get_load_from_state_dict_modulesitems)r1prefixlocal_metadatar child error_msgsloadmetadata missing_keys state_dictunexpected_keysrrrH`s z1BertPreTrainedModel.from_pretrained..loadbertcss|]}|dVqdS)bert.N) startswith).0srrr jsz6BertPreTrainedModel.from_pretrained..rN)rC)r<)r r,cuda is_availablerHrkeysreplaceappendzippopgetattrcopyr;hasattrany)clsnum_language_layersnum_vision_layersnum_mixed_layers bert_configrKr5Zfrom_tfpath_or_repo_idfilenamerrZ weights_pathmodelr7Zold_keysnew_keyskeynew_keyold_keyZ start_prefixrrFrfrom_pretrained5sD          z#BertPreTrainedModel.from_pretrainedrr)NNFr3r4) __name__ __module__ __qualname____doc__rr2 classmethodrj __classcell__rrrrr s r cs2eZdZdZd fdd 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*rrrrvs zBertAttention.__init__cCs6|dd|j|jf}|j|}|ddddS)Nr=r)sizerrruviewpermute)rxZ new_x_shaperrrtranspose_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)rwrgrxrr,matmul transposemathsqrtrur{rSoftmaxr contiguousrrvr) r hidden_statescontextZmixed_query_layerZmixed_key_layerZmixed_value_layerZ query_layerZ key_layerZ value_layerZattention_scoresZattention_probsZ context_layerZnew_context_layer_shaperrrforwards        zBertAttention.forwardrk)rlrmrnrorrrrqrrrrrps  rc*eZdZdZdfdd 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%rtdenser,r-ryhidden_dropout_probr{r|rrrrs zBertOutput.__init__cCs&||}||}|||}|Sr)rr{r-)rr input_tensorrrrrs  zBertOutput.forwardrkrlrmrnrorrrqrrrrrsrcr) rzyBERT cross attention layer. Based on: BERT (pytorch-transformer) https://github.com/huggingface/transformers rrcs6tt||_t||_t||_t||_dSr)rrratt_xroutput_xatt_youtput_yr|rrrrs    zBertMixedLayer.__init__cCs0|||}|||}||||||fSr)rrrr)rryrrrrrrs  zBertMixedLayer.forwardrkrrrrrrsrcr) rzpBERT pooler layer. Based on: BERT (pytorch-transformer) https://github.com/huggingface/transformers rrcs&tt|||_t|_dSr)rrrr%rTanh activation)rrtrrrrs zPooler.__init__cCs(|dddf}||}||}|SNr)rr)rrZfirst_token_tensorZ pooled_outputrrrrs  zPooler.forwardrkrrrrrrsrcs,eZdZdZdfd d Zdd dZZS)rz? Multimodal Transformers From Pretrained BERT Weights" r_rsr`rarbdictrrcsttdtf|_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. objcg|]}tjqSrr r*rP_rrr z'MultiModal.__init__..crrrrrrrrrcrr)rr*rrrrrrN)rrtypeobjectr*r embeddingsr ModuleListrangelanguage_encodervision_encoder mixed_encoderapplyr2)rr_r`rarbrrrrs  zMultiModal.__init__NcCsb|||}|jD] }||dd}q |jD] }|||d}q|jD] }|||\}}q#||fSr)rrrr)r input_idstoken_type_ids vision_featsattention_maskZlanguage_featureslayerrrrrs    zMultiModal.forward) r_rsr`rsrarsrbrrr)NNNrrrrrrsrcs^eZdZdZ                 dBdCfd=d> ZdDd@dAZZS)Erz 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:wr3r4 in_channelsrsimg_sizeSequence[int] | int patch_sizeint | tuple[int, int] num_classesr_r`rartdrop_outfloatrzgradient_checkpointingbool hidden_actstrrr+intermediate_sizelayer_norm_epsmax_position_embeddings model_typerrnum_hidden_layers pad_token_idposition_embedding_typetransformers_versiontype_vocab_size use_cache vocab_sizechunk_size_feed_forward is_decoderadd_cross_attentionrcstr | PathLikerdrrc !stid| ddd| d| d| d|d|d |d |d |d |d |d|d|d|d|d||||||dd} d| krPdksUtdtd|d|ddksi|d|ddkrmtdtj|||| ||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) rzZclassifier_dropoutNrrrrtr+rrrrrrrrrrreager)rrrrrZ_attn_implementationrr~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)rt)rr ValueErrorrrj multimodalr num_patchesrConv2d vision_projr-norm_vision_pos Parameterr,zeros pos_embed_visrpoolerrydropr%cls_head)!rrrrrr_r`rartrrzrrrr+rrrrrrrrrrrrrrrrrcrdrbrrrrs B     ( &   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_featureslogitsrrrrls    zTranschex.forward)rr rFrrrrrrrMrrrrrr}TrrFFr3r4)@rrsrrrrrrsr_rsr`rsrarsrtrsrrrzrrrrrrrr+rrrsrrrrsrrrrrsrrsrrsrrrrrrsrrrrsrrsrrrrrcrrdrrr)NNrrrrrrs8vr) __future__rrcollections.abcrr,rmonai.config.type_definitionsr monai.utilsrrrr r r __all__Moduler rrrrrrrrrrs(     P&"