a
    d6                     @   s   d Z ddlmZmZmZmZ ddlZddlmZ ddl	m  m
Z ddl	mZmZ G dd dejZG dd dejZG d	d
 d
ejZG dd dejZG dd dejZdd ZG dd dejZG dd dejZG dd dejZG dd dejZdS )z%Implements several backbone networks.    )DictListTupleUnionN)pixel_shufflesoftmaxc                       sD   e Zd ZdZdeeeeed fddZejejd	d
dZ  Z	S )HourglassBackbonea  Hourglass network, taken from https://github.com/zhou13/lcnn.

    Args:
        input_channel: number of input channels.
        depth: number of residual blocks per hourglass module.
        num_stacks: number of hourglass modules stacked together.
        num_blocks: number of layers in each residual block.
        num_classes: number of heads for the output of a hourglass module.
                )input_channeldepth
num_stacks
num_blocksnum_classesc              
      s4   t    t| _tf i | j|||||d| _d S )Nheadr   r   r   r   input_channels)super__init__MultitaskHeadr   hgnet)selfr   r   r   r   r   	__class__ g/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/kornia/feature/sold2/backbones.pyr      s    
zHourglassBackbone.__init__)input_imagesreturnc                 C   s
   |  |S N)r   )r   r   r   r   r   forward'   s    zHourglassBackbone.forward)r	   r
   r   r	   r   
__name__
__module____qualname____doc__intr   torchTensorr"   __classcell__r   r   r   r   r      s    
r   c                       s6   e Zd Zed fddZejejdddZ  ZS )r   )r   c                    s   t    t|d }dgdgdgg}g }t|g D ]:}|ttj||dddtjddtj||dd q4t	|| _
d S )	Nr
   r   r	      )kernel_sizepaddingTZinplacer-   )r   r   r(   sumappendnn
SequentialConv2dReLU
ModuleListheads)r   r   mZ	head_sizer8   Zoutput_channelsr   r   r   r   ,   s    

zMultitaskHead.__init__xr    c                    s   t j fdd| jD ddS )Nc                    s   g | ]}| qS r   r   ).0r   r;   r   r   
<listcomp>=       z)MultitaskHead.forward.<locals>.<listcomp>r	   Zdim)r)   catr8   r   r;   r   r=   r   r"   <   s    zMultitaskHead.forward)	r$   r%   r&   r(   r   r)   r*   r"   r+   r   r   r   r   r   +   s   r   c                       sR   e Zd Zd	eeeeeeef f ejjd fddZ	ej
ej
dddZ  ZS )
Bottleneck2Dr	   N)inplanesplanesstride
downsamplec                    s   t    t|| _tj||dd| _t|| _tj||d|dd| _t|| _	tj||d dd| _
tjdd| _|| _|| _d S )Nr	   r0   r,   r-   rF   r.   r   Tr/   )r   r   r3   BatchNorm2dbn1r5   conv1bn2conv2bn3conv3r6   relurG   rF   )r   rD   rE   rF   rG   r   r   r   r   A   s    
zBottleneck2D.__init__r:   c                 C   s~   |}|  |}| |}| |}| |}| |}| |}| |}| |}| |}| jd urr| |}||7 }|S r!   )rJ   rP   rK   rL   rM   rN   rO   rG   )r   r;   Zresidualoutr   r   r   r"   P   s    










zBottleneck2D.forward)r	   N)r$   r%   r&   r(   r   r   r)   r3   Moduler   r*   r"   r+   r   r   r   r   rC   @   s
    rC   c                       s   e Zd Zdejjeeeed fddZejjeeejjdddZejjeeeejj	dd	d
Z
eejejdddZejejdddZ  ZS )	Hourglassr   )blockr   rE   r   	expansionc                    s2   t    || _|| _|| _| ||||| _d S r!   )r   r   r   rT   rU   _make_hour_glassr   )r   rT   r   rE   r   rU   r   r   r   r   h   s
    
zHourglass.__init__)rT   r   rE   r    c                 C   s4   g }t d|D ]}|||| j | qtj| S )Nr   )ranger2   rU   r3   r4   )r   rT   r   rE   layers_r   r   r   _make_residualo   s    zHourglass._make_residual)rT   r   rE   r   r    c           	   	   C   sn   g }t |D ]V}g }t dD ]}|| ||| q|dkrR|| ||| |t| qt|S )Nr,   r   )rW   r2   rZ   r3   r7   )	r   rT   r   rE   r   hgliresrY   r   r   r   rV   u   s    zHourglass._make_hour_glass)nr;   r    c           	      C   s   | j |d  d |}tj|ddd}| j |d  d |}|dkrV| |d |}n| j |d  d |}| j |d  d |}tj||jdd  d}|| }|S )Nr	   r   r   rF   r,   )size)r   FZ
max_pool2d_hour_glass_forwardZinterpolateshape)	r   r^   r;   Zup1Zlow1Zlow2Zlow3Zup2rQ   r   r   r   rb      s    zHourglass._hour_glass_forwardr:   c                 C   s   |  | j|S r!   )rb   r   rB   r   r   r   r"      s    zHourglass.forward)r   )r$   r%   r&   r)   r3   rR   r(   r   rZ   r7   rV   r*   rb   r"   r+   r   r   r   r   rS   g   s
    rS   c                
       s   e Zd ZdZdejjejjeeeeeed fddZdejjeee	ee
eef f ejjddd	Zeeejjd
ddZejejdddZ  ZS )HourglassNetz,Hourglass model from Newell et al ECCV 2016.r   )rT   r   r   r   r   r   r   rU   c	                    s  t    d| _d| _|| _|| _tj|| jdddd| _t	| j| _
tjdd| _| || jd	| _| || jd	| _| || jd	| _tjddd
| _| j| j }	g g g g g g f\}
}}}}}t|D ]}|
t||| j| || || j| || |	|	 |||	 ||d	 k r|tj|	|	d	d |tj||	d	d qt|
| _t|| _t|| _t|| _t|| _t|| _d S )N@         r   r,   rH   Tr/   r	   r_   r0   )r   r   rD   Z	num_featsr   rU   r3   r5   rK   rI   rJ   r6   rP   rZ   layer1layer2layer3Z	MaxPool2dmaxpoolrW   r2   rS   _make_fcr7   r   r]   fcscorefc_score_)r   rT   r   r   r   r   r   r   rU   chr[   r]   rm   rn   ro   rp   r\   r   r   r   r      s8    
zHourglassNet.__init__r	   )rT   rE   blocksrF   r    c                 C   s   d }|dks| j || j kr<ttj| j || j d|d}g }||| j ||| || j | _ td|D ]}||| j | qltj| S )Nr	   )r-   rF   )rD   rU   r3   r4   r5   r2   rW   )r   rT   rE   rr   rF   rG   rX   rY   r   r   r   rZ      s     zHourglassNet._make_residual)rD   	outplanesr    c                 C   s*   t |}t j||dd}t ||| jS )Nr	   r0   )r3   rI   r5   r4   rP   )r   rD   rs   Zbnconvr   r   r   rl      s    
zHourglassNet._make_fcr:   c                 C   s   g }|  |}| |}| |}| |}| |}| |}| |}t| jD ]|}| j	| |}| j
| |}| j| |}| j| |}|| || jd k rT| j| |}| j| |}|| | }qT|S Nr	   )rK   rJ   rP   rh   rk   ri   rj   rW   r   r   r]   rm   rn   r2   ro   rp   )r   r;   rQ   r\   yrn   ro   rp   r   r   r   r"      s&    







zHourglassNet.forward)r   )r	   )r$   r%   r&   r'   r)   r3   rR   r(   r   r   r   rZ   rl   r*   r"   r+   r   r   r   r   rd      s&    , rd   c               	   K   s:   t t| ddd | d | d | d | d | d d	}|S )
Nr   c                 S   s   t | |dS ru   )r3   ZConv2D)Zc_inZc_outr   r   r   <lambda>   r?   zhg.<locals>.<lambda>r   r   r   r   r   r   )rd   rC   get)kwargsmodelr   r   r   r      s    	r   c                       s>   e Zd ZdZd
eed fddZejejddd	Z  Z	S )SuperpointDecoderzJunction decoder based on the SuperPoint architecture.

    Args:
        input_feat_dim: channel size of the input features.
    Returns:
        the junction heatmap, with shape (B, H, W).
    rf      )input_feat_dim	grid_sizec                    sT   t    tjjdd| _tjj|ddddd| _tjjddddd	d| _|| _	d S )
NTr/      r,   r   r	   rH   A   r   )
r   r   r)   r3   r6   rP   r5   convPaconvPbr~   )r   r}   r~   r   r   r   r     s
    
zSuperpointDecoder.__init__input_featuresr    c                 C   s^   |  | |}| |}t|dd}t|d d d dd d d d f | jd d df }|S )Nr	   r@   r   )rP   r   r   r   r   r~   )r   r   featsemiZ	junc_probZ	junc_predr   r   r   r"     s
    
4zSuperpointDecoder.forward)rf   r|   r#   r   r   r   r   r{      s   r{   c                       sT   e Zd ZdZdeeed fddZeee ddd	Zej	ej	d
ddZ
  ZS )PixelShuffleDecodera"  Pixel shuffle decoder used to predict the line heatmap.

    Args:
        input_feat_dim: channel size of the input features.
        num_upsample: how many upsamples are performed.
        output_channel: number of output channels.
    Returns:
        the (B, 1, H, W) line heatmap.
    rf   r   )r}   num_upsampleoutput_channelc                    s   t    | || _td| _g }|ttj	|| jd ddddt
| jd tjdd | jdd D ]6}|ttj	||ddddt
|tjdd qr|ttj	| jd |dddd t|| _d S )	Nr   r   r,   r	   rH   Tr/   r   )r   r   get_channel_confZchannel_confr3   ZPixelShuffle
pixshuffler2   r4   r5   rI   r6   r7   conv_block_lst)r   r}   r   r   r   Zchannelr   r   r   r   !  s.    

	
	zPixelShuffleDecoder.__init__)r   r    c                 C   s   |dkrg dS g dS )z2Get num of channels based on number of upsampling.r   )r   re      )r   re   r   r
   r   )r   r   r   r   r   r   D  s    z$PixelShuffleDecoder.get_channel_confr   c                 C   s`   |}| j d d D ]}||}| |}q| j d |}t|ddd d dd d d d f }|S )Nr   r	   r@   )r   r   r   )r   r   rQ   rT   heatmapr   r   r   r"   J  s    $zPixelShuffleDecoder.forward)rf   r   r   )r$   r%   r&   r'   r(   r   r   r   r)   r*   r"   r+   r   r   r   r   r     s   
#r   c                       s<   e Zd ZdZd	ed fddZejejdddZ  Z	S )
SuperpointDescriptorzDescriptor decoder based on the SuperPoint arcihtecture.

    Args:
        input_feat_dim: channel size of the input features.
    Returns:
        the semi-dense descriptors with shape (B, 128, H/4, W/4).
    rf   )r}   c                    sN   t    tjjdd| _tjj|ddddd| _tjjdddddd| _d S )	NTr/   r   r,   r	   rH   rf   r   )	r   r   r)   r3   r6   rP   r5   r   r   )r   r}   r   r   r   r   a  s    
zSuperpointDescriptor.__init__r   c                 C   s   |  | |}| |}|S r!   )rP   r   r   )r   r   r   r   r   r   r   r"   g  s    
zSuperpointDescriptor.forward)rf   r#   r   r   r   r   r   X  s   r   c                       s.   e Zd ZdZed fddZdd Z  ZS )SOLD2Netu  Full network for SOLD².

    Args:
        model_cfg: the configuration as a Dict.
    Returns:
        a Dict with the following values:
            junctions: heatmap of junctions.
            heatmap: line heatmap.
            descriptors: semi-dense descriptors.
    )	model_cfgc                    sb   t    || _tf i | jd | _d}t|| jd | _t|dd| _d| jv r^t	|| _
d S )NZbackbone_cfgr   r~   r   )r   use_descriptor)r   r   cfgr   backbone_netr{   junction_decoderr   heatmap_decoderr   descriptor_decoder)r   r   Zfeat_channelr   r   r   r   }  s    

zSOLD2Net.__init__c                 C   sD   |  |}| |}| |}||d}d| jv r@| ||d< |S )N)	junctionsr   r   Zdescriptors)r   r   r   r   r   )r   r   featuresr   Zheatmapsoutputsr   r   r   r"     s    




zSOLD2Net.forward)r$   r%   r&   r'   r   r   r"   r+   r   r   r   r   r   q  s   r   )r'   typingr   r   r   r   r)   Ztorch.nnr3   Ztorch.nn.functionalZ
functionalra   r   r   rR   r   r   rC   rS   rd   r   r{   r   r   r   r   r   r   r   <module>   s   '+[B