a
    
d                     @   s   d dl Z d dlZd dlmZ d dlmZ d dlmZ d dlm	Z	 d dl
mZ d dlmZmZmZ ddd	d	d
ZdddZdd Zdd ZG dd dejZdS )    N)Compose)DPTDepthModel)MidasNet)MidasNet_small)ResizeNormalizeImagePrepareForNetz(midas_models/dpt_large-midas-2f21e586.ptz)midas_models/dpt_hybrid-midas-501f0c75.pt 	dpt_large
dpt_hybrid	midas_v21midas_v21_smallTc                 C   s   | S )zbOverwrite model.train with this function to make sure train/eval mode
    does not change anymore. )selfmoder   r   h/var/www/html/stable-diffusion-webui/repositories/stable-diffusion-stability-ai/ldm/modules/midas/api.pydisabled_train   s    r   c              
   C   s   | dkr*d\}}d}t g dg dd}n| dkrTd\}}d}t g dg dd}nh| dkr~d\}}d}t g d	g d
d}n>| dkrd\}}d}t g d	g d
d}ndsJ d|  dtt||d dd|tjd|t g}|S )Nr     r   minimal      ?r   r   meanstdr   r   upper_boundg
ףp=
?gv/?gCl?gZd;O?gy&1?g?r      r    Fmodel_type '*' not implemented, use: --model_type largeT    Zresize_targetZkeep_aspect_ratioZensure_multiple_ofZresize_methodZimage_interpolation_method)r   r   r   cv2INTER_CUBICr   )
model_typenet_wnet_hresize_modenormalization	transformr   r   r   load_midas_transform   s@    	r-   c              
   C   s:  t |  }| dkr@t|ddd}d\}}d}tg dg dd}n| d	krxt|d
dd}d\}}d}tg dg dd}n| dkrt|dd}d\}}d}tg dg dd}n\| dkrt|ddddddid}d\}}d}tg dg dd}ntd|  d ds
J tt||d dd|tj	d|t
 g}| |fS )Nr   Z
vitl16_384T)pathbackbonenon_negativer   r   r   r   r   Zvitb_rn50_384r   )r0   r   r   r   r   @   efficientnet_lite3expand)featuresr/   
exportabler0   blocksr   r!   r"   Fr#   r$   )	ISL_PATHSr   r   r   r   printr   r   r%   r&   r   eval)r'   
model_pathmodelr(   r)   r*   r+   r,   r   r   r   
load_modelI   sh    

	r<   c                       s4   e Zd Zg dZg dZ fddZdd Z  ZS )MiDaSInference)Z	DPT_LargeZ
DPT_HybridZMiDaS_smallr
   c                    s6   t    || jv sJ t|\}}|| _t| j_d S )N)super__init__MODEL_TYPES_ISLr<   r;   r   train)r   r'   r;   _	__class__r   r   r?      s
    
zMiDaSInference.__init__c                 C   s   t  @ | |}t jjj|d|jdd  ddd}W d    n1 sN0    Y  |j|jd d|jd |jd fksJ |S )N      bicubicF)sizer   align_cornersr      )torchno_gradr;   nn
functionalinterpolate	unsqueezeshape)r   x
predictionr   r   r   forward   s    

$(zMiDaSInference.forward)__name__
__module____qualname__ZMODEL_TYPES_TORCH_HUBr@   r?   rT   __classcell__r   r   rC   r   r=      s   r=   )T)r%   rK   torch.nnrM   Ztorchvision.transformsr   Z!ldm.modules.midas.midas.dpt_depthr   Z!ldm.modules.midas.midas.midas_netr   Z(ldm.modules.midas.midas.midas_net_customr   Z"ldm.modules.midas.midas.transformsr   r   r   r7   r   r-   r<   Moduler=   r   r   r   r   <module>   s    
-@