o r) updatezipnpmaximumlenarcsinhsqrtshapezerosfloat32clip)imgsbandsscalesmZ rgbscalesIimgbandplanescaleQZfIHWrgbr*N/home/dell461/ljm/DATA_other/AstroM3/ExampleCode/example1/model/DesiEncoder.pysdss_rgbs0       (r,cKst||tddddddS)N)rg@)rg333333 @)rg@)r r rgQ?)rr )r,dict)Zrimgsrignoredr*r*r+dr2_rgb8sr/FcCsztj|tjd}tj|tjd}t||}tj|dd}|dd||g}t||}|r;tjtd|g|gdd}|S)z grid_size: int of the grid height and width return: pos_embed: [grid_size*grid_size, embed_dim] or [1+grid_size*grid_size, embed_dim] (w/ or w/o cls_token) dtyperaxisrr) rarangermeshgridstackreshape!get_2d_sincos_pos_embed_from_grid concatenater) embed_dim grid_size cls_tokenZgrid_hZgrid_wgrid pos_embedr*r*r+get_2d_sincos_pos_embedBs  r?cCsJ|ddksJt|d|d}t|d|d}tj||gdd}|S)Nrrrr2)!get_1d_sincos_pos_embed_from_gridrr9)r:r=Zemb_hZemb_wembr*r*r+r8Ts r8cCs||ddksJtj|dtd}||d}dd|}|d}td||}t|}t|}tj||gd d }|S) z} embed_dim: output dimension for each position pos: a list of positions to be encoded: size (M,) out: (M, D) rrr0g@r i'zm,d->mdrr2)rr4floatr7einsumsincosr9)r:posomegaoutZemb_sinZemb_cosrAr*r*r+r@_s     r@c Cs d|vr|d}|jd}|jj}|jjd|}t|jd|d}t|d}||krtd||||f|ddd|f}|dd|df} | d|||dddd } tj j j | ||fd d d } | dd dd dd } tj || fdd } | |d<dSdSdS)Nr>rB?z(Position interpolate from %dx%d to %dx%drrrrbicubicF)sizemode align_cornersdim)r patch_embed num_patchesr>intprintr7permutetorchnn functional interpolateflattencat) modelZcheckpoint_modelZpos_embed_checkpointembedding_sizerSZnum_extra_tokensZ orig_sizenew_sizeZ extra_tokensZ pos_tokensZ new_pos_embedr*r*r+interpolate_pos_embedys(     r`c seZdZdZdddddddddd ejd d f fd d ZddZddZddZ ddZ ddZ ddZ ddZ ddZd'ddZd(d!d"Zd#d$Zd%d&ZZS))MaskedAutoEncoderViTz8 Masked Autoencoder with VisionTransformer backbone rig@Fr cs*tt||||_|jj}ttdd|_ tjtd|ddd|_ t fddt |D|_ |_tjdd|_ttdd|_tjtd|ddd|_t fddt |D|_|_tj|d |dd|_| |_| |_|dS) NrF) requires_gradc sg|] }tddqST)qkv_bias norm_layerr.0r )r: mlp_ratiorj num_headsr*r+ z1MaskedAutoEncoderViT.__init__..T)biasc sg|] }tddqSrhrkrl)decoder_embed_dimdecoder_num_headsrnrjr*r+rprqr)super__init__rrRrSrX ParameterrWrr<r> ModuleListrangeblocksnormLinear decoder_embed mask_tokendecoder_pos_embeddecoder_blocks decoder_norm decoder_pred norm_pix_losslambda_consistencyinitialize_weights)selfimg_size patch_sizeZin_chansr:depthrors decoder_depthrtrnrjrrrS __class__)rsrtr:rnrjror+rvs(    zMaskedAutoEncoderViT.__init__cCst|jjdt|jjddd}|jjt | dt|j jdt|jjddd}|j jt | d|jj jj}tjj||jddgtjjj|jddtjjj|jdd||jdS)NrBrKT)r<rr)std)r?r>rrTrRrSdatacopy_rW from_numpyrC unsqueezerprojweightrXinitxavier_uniform_viewnormal_r<r~apply _init_weights)rr>rwr*r*r+rs"" z'MaskedAutoEncoderViT.initialize_weightscCst|tjr'tjj|jt|tjr#|jdur%tj|jddSdSdSt|tj r?tj|jdtj|jddSdS)Nrr ) isinstancerXr|rWrrrrr constant_ LayerNorm)rr r*r*r+rs  z"MaskedAutoEncoderViT._init_weightsc Cs|j\}}}t|d|}tj|||jd}tj|dd}tj|dd} |ddd|f} tj|d| ddd|d} tj ||g|jd} d| ddd|f<tj| d| d} | | | fS)z Perform per-sample random masking by per-sample shuffling. Per-sample shuffling is done by argsort random noise. x: [N, L, D], sequence rdevicerPNrBrQindexr) rrTrWrandrargsortgatherrrepeatones) rx mask_ratioNLDlen_keepnoise ids_shuffle ids_restoreids_keepx_maskedmaskr*r*r+random_maskings   z#MaskedAutoEncoderViT.random_maskingcCs|jjd}|jd|jdkr|jd|dksJ|jd|}}|j|jdd||||fd}td|}|j|jd|||ddfd}|S)zH imgs: (N, 3, H, W) x: (N, L, patch_size**2 *3) rrrrznchpwq->nhwpqc)rRrrr7rWrD)rrphrrr*r*r+patchifys * $zMaskedAutoEncoderViT.patchifycCs|jjd}t|jdd}}|||jdksJ|j|jd||||dfd}td|}|j|jdd||||fd}|S)zH x: (N, L, patch_size**2 *3) imgs: (N, 3, H, W) rrrKrrznhwpqc->nchpwq)rRrrTrr7rWrD)rrrrrrr*r*r+ unpatchifys  "zMaskedAutoEncoderViT.unpatchifycCs||}||jddddddf}|||\}}}|j|jddddddf}||jddd}tj||fdd}|jD]}||}qE| |}|||fS)NrrrBrP) rRr>rr<expandrrWr\rzr{)rrrrrr< cls_tokensblkr*r*r+forward_encoder s  "    z$MaskedAutoEncoderViT.forward_encoderc Cs||}|j|jd|jdd|jdd}tj|ddddddf|gdd}tj|d|ddd|jdd}tj|ddddddf|gdd}||j}|j D]}||}q]| |}| |}|ddddddf}|S)NrrrPrBrr) r}r~rrrWr\rrrrrr)rrrZ mask_tokensx_rr*r*r+forward_decoders *(&(     z$MaskedAutoEncoderViT.forward_decodercCsp||}|jr |jddd}|jddd}|||dd}||d}|jdd}|||}|S)zl imgs: [N, 3, H, W] pred: [N, L, p*p*3] mask: [N, L], 0 is keep, 1 is move. rBT)rQkeepdimrrKrrP)rrmeanvarsum)rrpredrtargetrrlossr*r*r+ forward_loss0s   z!MaskedAutoEncoderViT.forward_losscCs|r|dddf}|dddf}n|ddddfjdd}|ddddfjdd}tj|dd}tj|dd}dd||jdd}|S)u~ latent1, latent2: [N, L, D] use_cls: 是否使用 cls_token,如果 False 就用平均 patch 特征 NrrrPrBr)rF normalizer)rlatent1latent2Zuse_clsz1z2rr*r*r+consistency_lossCsz%MaskedAutoEncoderViT.consistency_loss?cCs|||\}}}|||}||||}|||\}} } ||| } ||| | } |||} || d|j| }|||fS)Nr)rrrrr)rrrrmask1Z ids_restore1Zpred1Z loss_recon1rmask2Z ids_restore2Zpred2Z loss_recon2Z loss_consZ loss_totalr*r*r+forwardUs    zMaskedAutoEncoderViT.forwardcCs<||}||jddddddf}|j\}}}tj|||jd}|}tj||d|dd}tj|dd} ||j dd } |ddd| f} tj |d| ddd|d} |j|jddddddf} | |jddd}tj|| fdd}|jD]}||}q||}||| fS)a Forward encoder using a given patch-level mask. Args: x: (N, 3, H, W) given_patch_mask: (N, L), 1 for masked patches, 0 for kept Returns: x: encoded tokens with cls token (N, len_keep + 1, embed_dim) mask: (N, L), same as input ids_restore: (N, L), mapping for unshuffling NrrrPrBrr)rRr>rrWrrrCrmaxrrTitemrrrr<rr\rzr{)rrgiven_patch_maskrrrrZ mask_floatrrrrrr<rrr*r*r+forward_encoder_with_given_maskis"    "    z4MaskedAutoEncoderViT.forward_encoder_with_given_maskcCs6|||\}}}|||}||||}|||fS)N)rrr)rrrZlatentrrrrr*r*r+forward_with_given_masks  z,MaskedAutoEncoderViT.forward_with_given_maskF)r)__name__ __module__ __qualname____doc__rXrrvrrrrrrrrrrrr __classcell__r*r*rr+ras& *   'racKs0td ddddddddttjddd |}|S) Nrci rerfr)eps) rr:rrorsrrtrnrjr*)rarrXr)kwargsr]r*r*r+mae_vit_base_patch16sr)Nrr)rWtorch.nnrX functoolsrZtimm.models.vision_transformerrrZtorch.nn.functionalrYrnumpyrZskimageZcv2cvr,r/r?r8r@r`Modulerarr*r*r*r+s.   )