a
    d?                     @   s  d Z ddlmZ ddlmZmZ ddlZddlmZ ddl	m
Z
mZ ddlmZmZmZ ddlmZmZmZ dd	lmZmZ dd
lmZmZ ddlmZmZ G dd dejZG dd dejZdWddZ dXddZ!dYddZ"ee"ddddde"ddddd dd!e"d"ddd#e"d$ddd dd%e" e"d&ddd d'e"d(ddd#e"d)ddd dd%e"d*dd+d,ddd-e"d.dd+d,dd/e"d0dd+d,d1e"d2dd+d,dd/e"e
ed3d4e"e
ed3d4e"e
ed3d4e"e
ed3d4d5Z#edZed6d7d8Z$ed[ed6d9d:Z%ed\ed6d;d<Z&ed]ed6d=d>Z'ed^ed6d?d@Z(ed_ed6dAdBZ)ed`ed6dCdDZ*edaed6dEdFZ+edbed6dGdHZ,edced6dIdJZ-edded6dKdLZ.edeed6dMdNZ/edfed6dOdPZ0ee1dQdRdSdSdTdUdV dS )ga   Hybrid Vision Transformer (ViT) in PyTorch

A PyTorch implement of the Hybrid Vision Transformers as described in:

'An Image Is Worth 16 x 16 Words: Transformers for Image Recognition at Scale'
    - https://arxiv.org/abs/2010.11929

`How to train your ViT? Data, Augmentation, and Regularization in Vision Transformers`
    - https://arxiv.org/abs/2106.10270

NOTE These hybrid model definitions depend on code in vision_transformer.py.
They were moved here to keep file sizes sane.

Hacked together by / Copyright 2020, Ross Wightman
    )partial)ListTupleN)IMAGENET_DEFAULT_MEANIMAGENET_DEFAULT_STD)StdConv2dSame	StdConv2d	to_2tuple   )generate_default_cfgsregister_modelregister_model_deprecations)	resnet26d	resnet50d)ResNetV2create_resnetv2_stem)_create_vision_transformerVisionTransformerc                       s*   e Zd ZdZd fdd		Zd
d Z  ZS )HybridEmbedd CNN Feature Map Embedding
    Extract feature map from CNN, flatten, project to embedding dim.
       r
   N      Tc              
      s  t    t|tjsJ t|}t|}|| _|| _|| _|d u rt	
 x |j}|r`|  | t	d||d |d }	t|	ttfr|	d }	|	jdd  }|	jd }
|| W d    n1 s0    Y  n.t|}t| jdr| jj d }
n| jj}
|d |d  dkr0|d |d  dks4J |d |d  |d |d  f| _| jd | jd  | _tj|
||||d| _d S )Nr
   r   feature_info)Zkernel_sizeZstridebias)super__init__
isinstancennModuler	   img_size
patch_sizebackbonetorchZno_gradtrainingevalzeroslisttupleshapeZtrainhasattrr   ZchannelsZnum_featuresZ	grid_sizeZnum_patchesZConv2dproj)selfr$   r"   r#   feature_sizein_chans	embed_dimr   r&   oZfeature_dim	__class__ n/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/timm/models/vision_transformer_hybrid.pyr   "   s4    



*0"zHybridEmbed.__init__c                 C   s@   |  |}t|ttfr |d }| |}|ddd}|S )Nr      r
   )r$   r   r)   r*   r-   flatten	transposer.   xr5   r5   r6   forwardJ   s    

zHybridEmbed.forward)r   r
   Nr   r   T)__name__
__module____qualname____doc__r   r<   __classcell__r5   r5   r3   r6   r      s         (r   c                       s>   e Zd ZdZd fdd		Zeejee	 f d
ddZ
  ZS )HybridEmbedWithSizer   r   r
   Nr   r   Tc              	      s   t  j|||||||d d S )N)r$   r"   r#   r/   r0   r1   r   )r   r   )r.   r$   r"   r#   r/   r0   r1   r   r3   r5   r6   r   W   s    
zHybridEmbedWithSize.__init__returnc                 C   sJ   |  |}t|ttfr |d }| |}|ddd|jdd  fS )Nr   r7   r
   r   )r$   r   r)   r*   r-   r8   r9   r+   r:   r5   r5   r6   r<   k   s
    

zHybridEmbedWithSize.forward)r   r
   Nr   r   T)r=   r>   r?   r@   r   r   r%   ZTensorr   intr<   rA   r5   r5   r3   r6   rB   S   s         rB   Fc                 K   s.   t t|d}|dd t| f||d|S )N)r$   r#   r
   )
pretrainedembed_layer)r   r   
setdefaultr   )variantr$   rF   kwargsrG   r5   r5   r6   !_create_vision_transformer_hybrids   s    rK   r      	   c              	   K   sx   | dd}|rdnd}|r(ttddn
ttdd}t| r\t| dd| dd	d
||d}nt| dd	|d
|d}|S )z ResNet-V2 backbone helperpadding_sameTZsame g:0yE>)Zepsr   r0   r   F)layersnum_classesZglobal_poolr0   preact	stem_type
conv_layer)rT   rS   rU   )getr   r   r   lenr   r   )rQ   rJ   rO   rT   rU   r$   r5   r5   r6   	_resnetv2y   s    rX   rP   c                 K   s    | ddd dddddddd	|S )
Ni  )r   r   r   ?ZbicubicT)      ?rZ   rZ   zpatch_embed.backbone.stem.convhead)urlrR   
input_sizeZ	pool_sizecrop_pctinterpolationZfixed_input_sizemeanstd
first_conv
classifierr5   )r\   rJ   r5   r5   r6   _cfg   s    rd   zhttps://storage.googleapis.com/vit_models/augreg/R_Ti_16-i21k-300ep-lr_0.001-aug_none-wd_0.03-do_0.0-sd_0.0--imagenet2012-steps_20k-lr_0.03-res_224.npzztimm/Tzpatch_embed.backbone.conv)r\   	hf_hub_idcustom_loadrb   zhttps://storage.googleapis.com/vit_models/augreg/R_Ti_16-i21k-300ep-lr_0.001-aug_none-wd_0.03-do_0.0-sd_0.0--imagenet2012-steps_20k-lr_0.03-res_384.npz)r     rg   g      ?)r\   re   rb   r]   r^   rf   zhttps://storage.googleapis.com/vit_models/augreg/R26_S_32-i21k-300ep-lr_0.001-aug_light0-wd_0.03-do_0.1-sd_0.1--imagenet2012-steps_20k-lr_0.03-res_224.npz)r\   re   rf   zhttps://storage.googleapis.com/vit_models/augreg/R26_S_32-i21k-300ep-lr_0.001-aug_medium2-wd_0.03-do_0.0-sd_0.0--imagenet2012-steps_20k-lr_0.03-res_384.npz)r\   re   r]   r^   rf   zthttps://github.com/rwightman/pytorch-image-models/releases/download/v0.1-vitjx/jx_vit_base_resnet50_384-9fd3c705.pth)r\   re   r]   r^   zhttps://storage.googleapis.com/vit_models/augreg/R50_L_32-i21k-300ep-lr_0.001-aug_medium1-wd_0.1-do_0.1-sd_0.1--imagenet2012-steps_20k-lr_0.01-res_224.npzzhttps://storage.googleapis.com/vit_models/augreg/R50_L_32-i21k-300ep-lr_0.001-aug_medium2-wd_0.1-do_0.0-sd_0.0--imagenet2012-steps_20k-lr_0.01-res_384.npzzohttps://storage.googleapis.com/vit_models/augreg/R_Ti_16-i21k-300ep-lr_0.001-aug_none-wd_0.03-do_0.0-sd_0.0.npziSU  rY   )r\   re   rR   r^   rb   rf   zshttps://storage.googleapis.com/vit_models/augreg/R26_S_32-i21k-300ep-lr_0.001-aug_medium2-wd_0.03-do_0.0-sd_0.0.npz)r\   re   rR   r^   rf   zzhttps://github.com/rwightman/pytorch-image-models/releases/download/v0.1-vitjx/jx_vit_base_resnet50_224_in21k-6f7c7740.pth)r\   re   rR   r^   zrhttps://storage.googleapis.com/vit_models/augreg/R50_L_32-i21k-300ep-lr_0.001-aug_medium2-wd_0.1-do_0.0-sd_0.0.npzzpatch_embed.backbone.conv1.0)r`   ra   rb   )z*vit_tiny_r_s16_p8_224.augreg_in21k_ft_in1kz*vit_tiny_r_s16_p8_384.augreg_in21k_ft_in1kz*vit_small_r26_s32_224.augreg_in21k_ft_in1kz*vit_small_r26_s32_384.augreg_in21k_ft_in1kzvit_base_r26_s32_224.untrained'vit_base_r50_s16_384.orig_in21k_ft_in1kz*vit_large_r50_s32_224.augreg_in21k_ft_in1kz*vit_large_r50_s32_384.augreg_in21k_ft_in1k"vit_tiny_r_s16_p8_224.augreg_in21k"vit_small_r26_s32_224.augreg_in21kvit_base_r50_s16_224.orig_in21k"vit_large_r50_s32_224.augreg_in21kz!vit_small_resnet26d_224.untrainedz%vit_small_resnet50d_s16_224.untrainedz vit_base_resnet26d_224.untrainedz vit_base_resnet50d_224.untrainedrC   c                 K   sF   t f ddi|}tddddd}td
|| d	t|fi |}|S )z3 R+ViT-Ti/S16 w/ 8x8 patch hybrid @ 224 x 224.
    rQ   r5            r   r#   r1   depth	num_headsvit_tiny_r_s16_p8_224r$   rF   )rs   rX   dictrK   rF   rJ   r$   Z
model_argsmodelr5   r5   r6   rs      s     rs   c                 K   sF   t f ddi|}tddddd}td
|| d	t|fi |}|S )z3 R+ViT-Ti/S16 w/ 8x8 patch hybrid @ 384 x 384.
    rQ   r5   rm   rn   ro   r   rp   vit_tiny_r_s16_p8_384rt   )ry   ru   rw   r5   r5   r6   ry      s     ry   c                 K   s@   t di |}tdddd}td	|| dt|fi |}|S )
 R26+ViT-S/S32 hybrid.
    r7   r7   r7   r7   rg   ro      r1   rq   rr   vit_small_r26_s32_224rt   )r{   )r~   ru   rw   r5   r5   r6   r~      s     r~   c                 K   s@   t di |}tdddd}td	|| dt|fi |}|S )
rz   r{   rg   ro   r|   r}   vit_small_r26_s32_384rt   )r{   )r   ru   rw   r5   r5   r6   r      s     r   c                 K   s@   t di |}tdddd}td|| dt|fi |}|S )	z R26+ViT-B/S32 hybrid.
    r{   r   ro   r}   vit_base_r26_s32_224rt   )r{   )r   ru   rw   r5   r5   r6   r     s     r   c                 K   s@   t di |}tdddd}td|| dt|fi |}|S )	zR R50+ViT-B/S16 hybrid from original paper (https://arxiv.org/abs/2010.11929).
    rL   r   ro   r}   vit_base_r50_s16_224rt   )rL   )r   ru   rw   r5   r5   r6   r     s     r   c                 K   s@   t di |}tdddd}td|| dt|fi |}|S )	z R50+ViT-B/16 hybrid from original paper (https://arxiv.org/abs/2010.11929).
    ImageNet-1k weights fine-tuned from in21k @ 384x384, source https://github.com/google-research/vision_transformer.
    rL   r   ro   r}   vit_base_r50_s16_384rt   )rL   )r   ru   rw   r5   r5   r6   r     s     r   c                 K   s@   t di |}tdddd}td	|| dt|fi |}|S )
 R50+ViT-L/S32 hybrid.
    r   rM   r|   r            r}   vit_large_r50_s32_224rt   )r   )r   ru   rw   r5   r5   r6   r   #  s     r   c                 K   s@   t di |}tdddd}td	|| dt|fi |}|S )
r   r   r   r   r   r}   vit_large_r50_s32_384rt   )r   )r   ru   rw   r5   r5   r6   r   .  s     r   c                 K   sN   t | |ddddgd}tddddd}td|| d
t|fi |}|S )zL Custom ViT small hybrid w/ ResNet26D stride 32. No pretrained weights.
    r0   r   TrM   rF   r0   Zfeatures_onlyZout_indicesr   rm   r1   rq   rr   Z	mlp_ratiovit_small_resnet26d_224rt   )r   r   rV   rv   rK   rw   r5   r5   r6   r   9  s     r   c                 K   sN   t | |ddddgd}tddddd}td
|| d	t|fi |}|S )zV Custom ViT small hybrid w/ ResNet50D 3-stages, stride 16. No pretrained weights.
    r0   r   Tr   r   rm   r   vit_small_resnet50d_s16_224rt   )r   r   rV   rv   rK   rw   r5   r5   r6   r   D  s     r   c                 K   sL   t | |ddddgd}tdddd}td|| d
t|fi |}|S )zK Custom ViT base hybrid w/ ResNet26D stride 32. No pretrained weights.
    r0   r   TrM   r   r   ro   r}   vit_base_resnet26d_224rt   )r   r   rw   r5   r5   r6   r   O  s     r   c                 K   sL   t | |ddddgd}tdddd}td|| d
t|fi |}|S )zK Custom ViT base hybrid w/ ResNet50D stride 32. No pretrained weights.
    r0   r   TrM   r   r   ro   r}   vit_base_resnet50d_224rt   )r   r   rw   r5   r5   r6   r   Z  s     r   ri   rj   rk   rl   rh   )Zvit_tiny_r_s16_p8_224_in21kZvit_small_r26_s32_224_in21kZvit_base_r50_s16_224_in21kZvit_base_resnet50_224_in21kZvit_large_r50_s32_224_in21kZvit_base_resnet50_384)F)rL   )rP   )F)F)F)F)F)F)F)F)F)F)F)F)F)2r@   	functoolsr   typingr   r   r%   Ztorch.nnr    Z	timm.datar   r   Ztimm.layersr   r   r	   	_registryr   r   r   Zresnetr   r   Zresnetv2r   r   Zvision_transformerr   r   r!   r   rB   rK   rX   rd   Zdefault_cfgsrs   ry   r~   r   r   r   r   r   r   r   r   r   r   r=   r5   r5   r5   r6   <module>   s   5 


B











