a
    d                     @   s   d dl mZmZmZ d dlZd dlmZ d dlmZm	Z	m
Z
mZ d dlmZ d dlmZ d dlmZ d dlmZ G dd	 d	ejZdS )
    )DictOptionalTupleN)DescriptorMatcherGFTTAffNetHardNetLocalFeatureMatcherLoFTR)LocalFeature)transform_points)RANSAC)warp_perspectivec                       s   e Zd ZdZdee eej eej edd fddZ	e
ejddd	Ze
ejdd
dZe ejddddZddddZeejef dddZejeejef dddZejeejef dddZejeejef dddZ  ZS )HomographyTrackera  Module, which performs local-feature-based tracking of the target planar object in the sequence of the
    frames.

    Args:
        initial_matcher: image matching module, e.g. :class:`~kornia.feature.LocalFeatureMatcher`
                          or :class:`~kornia.feature.LoFTR`. Default: :class:`~kornia.feature.GFTTAffNetHardNet`.
        fast_matcher: fast image matching module, e.g. :class:`~kornia.feature.LocalFeatureMatcher`
                          or :class:`~kornia.feature.LoFTR`. Default: :class:`~kornia.feature.DescriptorMatcher`.
        ransac: homography estimation module. Default: :class:`~kornia.geometry.RANSAC`.
        minimum_inliers_num: threshold for number inliers for matching to be successful.
    N   )initial_matcherfast_matcherransacminimum_inliers_numreturnc                    s   t    |p ttdtdd| _|p.td| _|pFtdddddd	| _	|| _
|  i | _i | _d | _d
| _d
| _d
| _|   d S )Ni  Zsmnngffffff?ZoutdoorZ
homographyg      @i   
   )Zinl_thZ
batch_sizeZmax_iterZmax_lo_itersr   )super__init__r   r   r   r   r   r   r   r   r   target_initial_representationtarget_fast_representationprevious_homographyinliers_numkeypoints0_numkeypoints1_numreset_tracking)selfr   r   r   r   	__class__ g/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/kornia/tracking/planar_tracker.pyr      s    
zHomographyTracker.__init__)r   c                 C   s   | j jS N)targetdevicer   r!   r!   r"   r%   5   s    zHomographyTracker.devicec                 C   s   | j jS r#   )r$   dtyper&   r!   r!   r"   r'   9   s    zHomographyTracker.dtype)r$   r   c                 C   sJ   || _ i | _i | _t| jdr,| j|| _t| jdrF| j|| _d S )Nextract_features)r$   r   r   hasattrr   r(   r   )r   r$   r!   r!   r"   
set_target=   s    zHomographyTracker.set_targetc                 C   s
   d | _ d S r#   )r   r&   r!   r!   r"   r   G   s    z HomographyTracker.reset_trackingc                 C   s,   d| _ d| _d| _tjdd| j| jddfS )Nr      )r%   r'   F)r   r   r   torchemptyr%   r'   r&   r!   r!   r"   no_matchJ   s    zHomographyTracker.no_match)xr   c           
      C   s   | j |d}| j D ]\}}||| d< q| |}|d |d dk }|d |d dk }t|| _t|| _| j| jk r|  S | 	||\}}	|	
  | _| j| jk r|  S | | _|dfS )zHThe frame `x` is matched with initial_matcher and  verified with ransac.Zimage0Zimage10
keypoints0batch_indexesr   
keypoints1T)r$   r   itemsr   lenr   r   r   r.   r   sumitemr   cloner   )
r   r/   
input_dictkv
match_dictr2   r4   Hinliersr!   r!   r"   match_initialP   s     



zHomographyTracker.match_initialc                 C   s  | j dur| j  d }|ddddddf d |ddddddf< |dddddf  d8  < t|}| jjdd \}}t||||f}| j|d}| j D ]\}}	|	|| d< q| 	|}
|
d |
d	 dk }|
d
 |
d	 dk }t
||}t|| _t|| _| j| jk r4|   |  S | ||\}}|  | _| j| jk rp|   |  S | | _ |dfS )z~The frame `x` is prewarped according to the previous frame homography, matched with fast_matcher
        verified with ransac.Nr      g?g      $@r0   r1   r2   r3   r4   T)r   r9   r,   Zinverser$   shaper   r   r5   r   r
   r6   r   r   r   r   r.   r   r7   r8   r   )r   r/   ZHwarpZHinvhwZframe_warpedr:   r;   r<   r=   r2   r4   r>   r?   r!   r!   r"   track_next_framej   s4    
4





z"HomographyTracker.track_next_framec                 C   s   | j d ur| |S | |S r#   )r   rE   r@   )r   r/   r!   r!   r"   forward   s    

zHomographyTracker.forward)NNNr   )__name__
__module____qualname____doc__r   r	   nnModuleintr   propertyr,   r%   r'   Zno_gradZTensorr*   r   r   boolr.   r@   rE   rF   __classcell__r!   r!   r   r"   r      s.       	%r   )typingr   r   r   r,   Ztorch.nnrK   Zkornia.featurer   r   r   r   Zkornia.feature.integratedr	   Zkornia.geometry.linalgr
   Zkornia.geometry.ransacr   Zkornia.geometry.transformr   rL   r   r!   r!   r!   r"   <module>   s   