a
    dF                     @   s8  d dl mZmZmZmZm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gZejed	d
dZd eej eeejee f dddZejejejdddZejejejejejdddZd!ejeeejdddZejejejejejejejdddZG dd dZejjG dd dZdS )"    )ListOptionalTupleUnioncastN)Tensor)validate_bbox)transform_points)eye_likeBoxesBoxes3D)dtypereturnc                 C   s   | t jt jt jt jt jfv S N)torchfloat16float32float64Zbfloat16Zhalf)r    r   ^/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/kornia/geometry/boxes.py_is_floating_point_dtype   s    r   pad)boxesmethodr   c                    s   t dd | D s,tddd | D  d|dkrntdd | D   fd	d| D }tjjjj| d
d}ntd| d||fS )z&Merge a list of boxes into one tensor.c                 s   s6   | ].}|j d d tddgko,| dkV  qdS )N         )shaper   Sizedim.0Zboxr   r   r   	<genexpr>       z"_merge_box_list.<locals>.<genexpr>z5Input boxes must be a list of (N, 4, 2) shaped. Got: c                 S   s   g | ]
}|j qS r   r   r!   r   r   r   
<listcomp>   r$   z#_merge_box_list.<locals>.<listcomp>.r   c                 s   s   | ]}|j d  V  qdS )r   Nr%   r!   r   r   r   r#      r$   c                    s   g | ]} |j d   qS )r   r%   r!   Zmax_Nr   r   r&      r$   T)Zbatch_first`z` is not implemented.)	all	TypeErrormaxr   nnutilsZrnnZpad_sequenceNotImplementedError)r   r   statsoutputr   r(   r   _merge_box_list   s    r2   )r   Mr   c                 C   s   |  r|n| }| jdd \}}}| d|| |}|jdkrH|n|d}|jd |jd krtd|jd  d|jd  dt||}|| }|S )	a  Transforms 3D and 2D in kornia format by applying the transformation matrix M. Boxes and the transformation
    matrix could be batched or not.

    Args:
        boxes: 2D quadrilaterals or 3D hexahedrons in kornia format.
        M: the transformation matrix of shape :math:`(3, 3)` or :math:`(B, 3, 3)` for 2D and :math:`(4, 4)` or
            :math:`(B, 4, 4)` for 3D hexahedron.
    Nr   r   zBatch size mismatch. Got z for boxes and z for the transformation matrix.)	is_floating_pointfloatr   viewndim	unsqueeze
ValueErrorr	   Zview_as)r   r3   Zboxes_per_batchZn_points_per_boxZcoordinates_dimensionZpointstransformed_boxesr   r   r   _transform_boxes    s    	

r=   )xminyminwidthheightr   c                 C   s   | j |j   kr0|j   kr0|j   kr0dks:n tdtj| jd | jd ddf| j| jd}| d|d< |d|d	< |d
  |d 7  < |d  |d 7  < |d  |d 7  < |d  |d 7  < |S )Nr   zXWe expect to create a batch of 2D boxes (quadrilaterals) in vertices format (B, N, 4, 2)r      r   devicer   r5   .r   .rB   .rB   r   .r   r   .r   rB   .r   rB   )r9   r;   r   zerosr   rD   r   r:   )r>   r?   r@   rA   Zpolygonsr   r   r   _boxes_to_polygons:   s    2(rL   xyxyTr   modevalidate_boxesr   c           	      C   s  |  }|drl| jdk}d| j  kr2dkrRn n| jdd tddgkstd| d| j d	nd|d
r| jdk}d| j  krdkrn n| jd dkstd| d| j d	ntd| |  r| n|  } |r| n| 	d} |dr|dkr^| 
 }|ddddf d |ddddf< |ddddf d |ddddf< n"|dkrr| 
 }ntd| | pt|  n|d
r|dkr| d | d  | d | d   }}nb|dkr| d | d  d | d | d  d  }}n,|dkr"| d | d  }}ntd| |rb|dk rLtd|dk rbtd| d | d  }}t||||}ntd| |r|n|d}|S )z%Convert from boxes to quadrilaterals.verticesr   r   r   Nr   z3Boxes shape must be (N, 4, 2) or (B, N, 4, 2) when z mode. Got r'   Zxyr5   z-Boxes shape must be (N, 4) or (B, N, 4) when Unknown mode r   .rB   vertices_plusrM   .r   rF   .r   rE   	xyxy_plusxywh%Some boxes have negative widths or 0.&Some boxes have negative heights or 0.)lower
startswithr9   r   r   r   r;   r6   r7   r:   cloner   anyrL   squeeze)	r   rO   rP   batchedquadrilateralsrA   r@   r>   r?   r   r   r   _boxes_to_quadrilateralsL   sN    

6

(
$&


$
,
ra   )r>   r?   zminr@   rA   depthr   c           	      C   s,  | j |j   krH|j   krH|j   krH|j   krH|j   krHdksRn tdtj| jd | jd ddf| j| jd}| d|d	< |d|d
< |d|d< |d  |d 7  < |d  |d 7  < |d  |d 7  < |d  |d 7  < | }|d  |dd 7  < tj	||gdd}|S )Nr   zUWe expect to create a batch of 3D boxes (hexahedrons) in vertices format (B, N, 8, 3)r   rB   r   r   rC   r5   rE   rF   rU   rG   rH   rI   rJ   r   r    )
r9   r;   r   rK   r   rD   r   r:   r\   cat)	r>   r?   rb   r@   rA   rc   Zfront_verticesZback_verticesZ
polygons3dr   r   r   _boxes3d_to_polygons3d   s    J(rf   c                   @   s  e Zd ZdZd;eejeej f ee	ddddZ
eejejf dd	d
Zd<d ed dddZd=eeeeeef f  eeeeeef f  ed dddZd>eed dddZd?ee ee ed dddZejdddZed@eejeej f e	ed dddZdAee	 eeejeej f ddd Zeeejd!d"d#ZdBejed d$d%d&Zejd d'd(d)ZdCee	ed d+d,d-Zeejdd.d/Zee	dd0d1Zeejdd2d3Zeej dd4d5Z dDeej eej  d d6d7d8Z!d dd9d:Z"dS )Er   a   2D boxes containing N or BxN boxes.

    Args:
        boxes: 2D boxes, shape of :math:`(N, 4, 2)`, :math:`(B, N, 4, 2)` or a list of :math:`(N, 4, 2)`.
            See below for more details.
        raise_if_not_floating_point: flag to control floating point casting behaviour when `boxes` is not a floating
            point tensor. True to raise an error when `boxes` isn't a floating point tensor, False to cast to float.
        mode: the box format of the input boxes.

    Note:
        **2D boxes format** is defined as a floating data type tensor of shape ``Nx4x2`` or ``BxNx4x2``
        where each box is a `quadrilateral <https://en.wikipedia.org/wiki/Quadrilateral>`_ defined by it's 4 vertices
        coordinates (A, B, C, D). Coordinates must be in ``x, y`` order. The height and width of a box is defined as
        ``width = xmax - xmin + 1`` and ``height = ymax - ymin + 1``. Examples of
        `quadrilaterals <https://en.wikipedia.org/wiki/Quadrilateral>`_ are rectangles, rhombus and trapezoids.
    TrS   Nr   raise_if_not_floating_pointrO   r   c                 C   s   d | _ t|trt|\}| _ t|tjs>tdt| d| sb|rZt	d|j
 | }t|jdkrz|d}d|j  krdkrn n|jdd  d	kst	d
|j d|jdkrdnd| _|| _|| _d S )N"Input boxes is not a Tensor. Got: r'   +Coordinates must be in floating point. Got r   )r5   r   r   r   r   )r   r   z3Boxes shape must be (N, 4, 2) or (B, N, 4, 2). Got FT)_N
isinstancelistr2   r   r   r+   typer6   r;   r   r7   lenr   reshaper9   _is_batched_data_modeselfr   rh   rO   r   r   r   __init__   s     

,zBoxes.__init__r   c                 C   s0   t tj| jddd}|d |d  }}||fS )a  Compute boxes heights and widths.

        Returns:
            - Boxes heights, shape of :math:`(N,)` or :math:`(B,N)`.
            - Boxes widths, shape of :math:`(N,)` or :math:`(B,N)`.

        Example:
            >>> boxes_xyxy = torch.tensor([[[1,1,2,2],[1,1,3,2]]])
            >>> boxes = Boxes.from_tensor(boxes_xyxy)
            >>> boxes.get_boxes_shape()
            (tensor([[1., 1.]]), tensor([[1., 2.]]))
        rW   Tas_padded_sequencerU   rT   )r   r   r   	to_tensor)ru   Z
boxes_xywhwidthsheightsr   r   r   get_boxes_shape   s    zBoxes.get_boxes_shapeF)r   inplacer   c                 C   s.   t j| j|jgdd}|r$|| _| S t|dS )a"  Merges boxes.

        Say, current instance holds :math:`(B, N, 4, 2)` and the incoming boxes holds :math:`(B, M, 4, 2)`,
        the merge results in :math:`(B, N + M, 4, 2)`.

        Args:
            boxes: 2D boxes.
            inplace: do transform in-place and return self.
        rB   rd   F)r   re   rr   datar   )ru   r   r~   r   r   r   r   merge   s
    
zBoxes.merge)topleftbotrightr~   r   c           	      C   sZ  t |trt |tst|r$| j}n
| j }|dddddf d|dd}||d |k  |d |d |k < |dddddf d|dd}||d |k  |d |d |k < |dddddf d|dd}||d |k |d |d |k< |dddddf d|dd}||d |k |d |d |k< |rP| S t|dS ) NrB   r   rE   rF   F)rl   r   r/   rr   r\   repeatsizer   )	ru   r   r   r~   rr   Z	topleft_xZ	topleft_yZ
botright_xZ
botright_yr   r   r   clamp   s     
& & & & zBoxes.clamp)correspondence_preserver~   r   c                 C   s   t dS )a  Trim out zero padded boxes.

        Given box arrangements of shape :math:`(4, 4, Box)`:
            -- Box -- Box -- Box  -- Box --
            --  0  --  0  -- Box  -- Box --
            --  0  -- Box --  0   --  0  --
            --  0  --  0  --  0   --  0  --

        Nothing will change if correspondence_preserve is True. Only pure zero layers will be
        removed, resulting in shape :math:`(4, 3, Box)`:
            -- Box -- Box -- Box  -- Box --
            --  0  --  0  -- Box  -- Box --
            --  0  -- Box --  0   --  0  --

        Otherwise, you will get :math:`(4, 2, Box)`:
            -- Box -- Box -- Box  -- Box --
            --  0  -- Box -- Box  -- Box --
        N)r/   )ru   r   r~   r   r   r   trim  s    z
Boxes.trim)min_areamax_arear~   r   c                 C   sX   |   }|r| j}n
| j }|d ur2d|||k < |d urFd|||k< |rN| S t|dS )Ng        F)compute_arearr   r\   r   )ru   r   r   r~   arearr   r   r   r   filter_boxes_by_area-  s    
zBoxes.filter_boxes_by_areac                 C   sp   | j ddddddf | j ddddddf  }| j ddddddf | j ddddddf  }|| S )zReturns :math:`(B, N)`.NrB   r   r   rr   )ru   whr   r   r   r   =  s    44zBoxes.compute_arearM   rN   c                    s<   t |tjrt| d}n fdd|D }| |d S )ay	  Helper method to easily create :class:`Boxes` from boxes stored in another format.

        Args:
            boxes: 2D boxes, shape of :math:`(N, 4)`, :math:`(B, N, 4)`, :math:`(N, 4, 2)` or :math:`(B, N, 4, 2)`.
            mode: The format in which the boxes are provided.

                * 'xyxy': boxes are assumed to be in the format ``xmin, ymin, xmax, ymax`` where ``width = xmax - xmin``
                  and ``height = ymax - ymin``. With shape :math:`(N, 4)`, :math:`(B, N, 4)`.
                * 'xyxy_plus': similar to 'xyxy' mode but where box width and length are defined as
                  ``width = xmax - xmin + 1`` and ``height = ymax - ymin + 1``.
                  With shape :math:`(N, 4)`, :math:`(B, N, 4)`.
                * 'xywh': boxes are assumed to be in the format ``xmin, ymin, width, height`` where
                  ``width = xmax - xmin`` and ``height = ymax - ymin``. With shape :math:`(N, 4)`, :math:`(B, N, 4)`.
                * 'vertices': boxes are defined by their vertices points in the following ``clockwise`` order:
                  *top-left, top-right, bottom-right, bottom-left*. Vertices coordinates are in (x,y) order. Finally,
                  box width and height are defined as ``width = xmax - xmin`` and ``height = ymax - ymin``.
                  With shape :math:`(N, 4, 2)` or :math:`(B, N, 4, 2)`.
                * 'vertices_plus': similar to 'vertices' mode but where box width and length are defined as
                  ``width = xmax - xmin + 1`` and ``height = ymax - ymin + 1``. ymin + 1``.
                  With shape :math:`(N, 4, 2)` or :math:`(B, N, 4, 2)`.

            validate_boxes: check if boxes are valid rectangles or not. Valid rectangles are those with width
                and height >= 1 (>= 2 when mode ends with '_plus' suffix).

        Returns:
            :class:`Boxes` class containing the original `boxes` in the format specified by ``mode``.

        Examples:
            >>> boxes_xyxy = torch.as_tensor([[0, 3, 1, 4], [5, 1, 8, 4]])
            >>> boxes = Boxes.from_tensor(boxes_xyxy, mode='xyxy')
            >>> boxes.data  # (2, 4, 2)
            tensor([[[0., 3.],
                     [0., 3.],
                     [0., 3.],
                     [0., 3.]],
            <BLANKLINE>
                    [[5., 1.],
                     [7., 1.],
                     [7., 3.],
                     [5., 3.]]])
        rO   rP   c                    s   g | ]}t | qS r   )ra   r!   r   r   r   r&   t  r$   z%Boxes.from_tensor.<locals>.<listcomp>F)rl   r   r   ra   )clsr   rO   rP   r`   r   r   r   from_tensorC  s    .zBoxes.from_tensor)rO   ry   r   c                 C   s\  | j r| jn
| jd}tj|jdd|jddgdd|jd |jd d}|du r^| j	}|
 }|dv rpnR|dv r|d	 |d
  d |d |d  d  }}||d< ||d	< ntd| |dv rtjg d|j|jd}|| }|drt|d |d
 |d |d	 }| jdurB|sBtdd t|| jD }n| j rN|n|d}|S )a^  Cast :class:`Boxes` to a tensor. ``mode`` controls which 2D boxes format should be use to represent
        boxes in the tensor.

        Args:
            mode: the output box format. It could be:

                * 'xyxy': boxes are defined as ``xmin, ymin, xmax, ymax`` where ``width = xmax - xmin`` and
                  ``height = ymax - ymin``.
                * 'xyxy_plus': similar to 'xyxy' mode but where box width and length are defined as
                  ``width = xmax - xmin + 1`` and ``height = ymax - ymin + 1``.
                * 'xywh': boxes are defined as ``xmin, ymin, width, height`` where ``width = xmax - xmin``
                  and ``height = ymax - ymin``.
                * 'vertices': boxes are defined by their vertices points in the following ``clockwise`` order:
                  *top-left, top-right, bottom-right, bottom-left*. Vertices coordinates are in (x,y) order. Finally,
                  box width and height are defined as ``width = xmax - xmin`` and ``height = ymax - ymin``.
                * 'vertices_plus': similar to 'vertices' mode but where box width and length are defined as
                  ``width = xmax - xmin + 1`` and ``height = ymax - ymin + 1``. ymin + 1``.
            as_padded_sequence: whether to keep the pads for a list of boxes. This parameter is only valid
                if the boxes are from a box list.

        Returns:
            Boxes tensor in the ``mode`` format. The shape depends with the ``mode`` value:

                * 'vertices' or 'verticies_plus': :math:`(N, 4, 2)` or :math:`(B, N, 4, 2)`.
                * Any other value: :math:`(N, 4)` or :math:`(B, N, 4)`.

        Examples:
            >>> boxes_xyxy = torch.as_tensor([[0, 3, 1, 4], [5, 1, 8, 4]])
            >>> boxes = Boxes.from_tensor(boxes_xyxy)
            >>> assert (boxes_xyxy == boxes.to_tensor(mode='xyxy')).all()
        r   r   rd   rB   r   N)rM   rV   )rW   rQ   rS   rT   rF   rU   rE   rR   )rM   rQ   )r   r   rB   rB   rC   rQ   c                 s   s>   | ]6\}}t jj|t|jd  ddg d| g V  qdS )rB   r   N)r   r-   Z
functionalr   ro   r   )r"   onr   r   r   r#     s   z"Boxes.to_tensor.<locals>.<genexpr>)rq   rr   r:   r   stackaminamaxr8   r   rO   rZ   r;   	as_tensorrD   r   r[   rL   rk   rm   zipr^   )ru   rO   ry   batched_boxesr   rA   r@   offsetr   r   r   rz   z  s2    ""*


zBoxes.to_tensor)rA   r@   r   c                 C   s  | j jrtd| jrDtj| j jd | j jd ||f| j| jd}n"tj| j jd ||f| j| jd}t	tj
| jddd}|dd	d	d
f d| |ddd	d
f d| t|d|||dd  D ],\}}d||d |d |d |d
 f< q|S )a  Convert 2D boxes to masks. Covered area is 1 and the remaining is 0.

        Args:
            height: height of the masked image/images.
            width: width of the masked image/images.

        Returns:
            the output mask tensor, shape of :math:`(N, width, height)` or :math:`(B,N, width, height)` and dtype of
            :func:`Boxes.dtype` (it can be any floating point dtype).

        Note:
            It is currently non-differentiable.

        Examples:
            >>> boxes = Boxes(torch.tensor([[  # Equivalent to boxes = Boxes.from_tensor([[1,1,4,3]])
            ...        [1., 1.],
            ...        [4., 1.],
            ...        [4., 3.],
            ...        [1., 3.],
            ...   ]]))  # 1x4x2
            >>> boxes.to_mask(5, 5)
            tensor([[[0., 0., 0., 0., 0.],
                     [0., 1., 1., 1., 1.],
                     [0., 1., 1., 1., 1.],
                     [0., 1., 1., 1., 1.],
                     [0., 0., 0., 0., 0.]]])
        cBoxes.to_tensor isn't differentiable. Please, create boxes from tensors with `requires_grad=False`.r   rB   r   rD   rM   Trx   .Nr   r5   r   r   )rr   requires_gradRuntimeErrorrq   r   rK   r   r   rD   r   r   rz   clamp_r   r8   roundint)ru   rA   r@   maskZclipped_boxes_xyxymask_channelZbox_xyxyr   r   r   to_mask  s    "",&zBoxes.to_maskr3   r~   r   c                 C   sb   d|j   krdkr,n n|jdd dkr>td|j dt| j|}|rX|| _| S t|dS )	a  Apply a transformation matrix to the 2D boxes.

        Args:
            M: The transformation matrix to be applied, shape of :math:`(3, 3)` or :math:`(B, 3, 3)`.
            inplace: do transform in-place and return self.

        Returns:
            The transformed boxes.
        r   r   r   N)r   r   zAThe transformation matrix shape must be (3, 3) or (B, 3, 3). Got r'   F)r9   r   r;   r=   rr   r   ru   r3   r~   r<   r   r   r   transform_boxes  s    
,zBoxes.transform_boxesr3   r   c                 C   s   | j |ddS )z0Inplace version of :func:`Boxes.transform_boxes`Tr~   r   ru   r3   r   r   r   transform_boxes_  s    zBoxes.transform_boxes_warp)r   r   r~   r   c                 C   sJ   |dkrt n|dkrnt td|}||dddddf< | j||dS )a#  Translates boxes by the provided size.

        Args:
            size: translate size for x, y direction, shape of :math:`(B, 2)`.
            method: "warp" or "fast".
            inplace: do transform in-place and return self.

        Returns:
            The transformed boxes.
        fastr   r   Nr   r   )r/   r
   r   )ru   r   r   r~   r3   r   r   r   	translate  s    
zBoxes.translatec                 C   s   | j S r   r   ru   r   r   r   r   '  s    z
Boxes.datac                 C   s   | j S r   rs   r   r   r   r   rO   +  s    z
Boxes.modec                 C   s   | j jS zReturns boxes device.rr   rD   r   r   r   r   rD   /  s    zBoxes.devicec                 C   s   | j jS zReturns boxes dtype.rr   r   r   r   r   r   r   4  s    zBoxes.dtyperD   r   r   c                 C   s.   |durt |std| jj||d| _| S z)Like :func:`torch.nn.Module.to()` method.NzBoxes must be in floating pointrC   r   r;   rr   toru   rD   r   r   r   r   r   9  s    zBoxes.toc                 C   s   t | j dS )NF)r   rr   r\   r   r   r   r   r\   A  s    zBoxes.clone)TrS   )F)NNF)FF)NNF)rM   T)NF)F)r   F)NN)#__name__
__module____qualname____doc__r   r   r   r   boolstrrv   r   r}   r   r   r   r   r   r7   r   r   classmethodr   rz   r   r   r   r   propertyr   rO   rD   r   r   r\   r   r   r   r   r      sh     "     7 H5 c                   @   s"  e Zd ZdZd'ejeeddddZe	ejejejf dd	d
Z
ed(ejeed dddZd)eejdddZeeeejdddZd*ejed dddZejd dddZeejdddZeedddZeejdd d!Zeejdd"d#Zd+eej eej d d$d%d&ZdS ),r   a  3D boxes containing N or BxN boxes.

    Args:
        boxes: 3D boxes, shape of :math:`(N,8,3)` or :math:`(B,N,8,3)`. See below for more details.
        raise_if_not_floating_point: flag to control floating point casting behaviour when `boxes` is not a floating
            point tensor. True to raise an error when `boxes` isn't a floating point tensor, False to cast to float.

    Note:
        **3D boxes format** is defined as a floating data type tensor of shape ``Nx8x3`` or ``BxNx8x3`` where each box
        is a `hexahedron <https://en.wikipedia.org/wiki/Hexahedron>`_ defined by it's 8 vertices coordinates.
        Coordinates must be in ``x, y, z`` order. The height, width and depth of a box is defined as
        ``width = xmax - xmin + 1``, ``height = ymax - ymin + 1`` and ``depth = zmax - zmin + 1``. Examples of
        `hexahedrons <https://en.wikipedia.org/wiki/Hexahedron>`_ are cubes and rhombohedrons.
    Txyzxyz_plusNrg   c                 C   s   t |tjs tdt| d| sF|r>td|j d| }t	|j
dkr^|d}d|j  krtdkrn n|j
dd  d	kstd
|j
 d|jdkrdnd| _|| _|| _d S )Nri   r'   rj   r   )r5      r   r   r   )   r   z53D bbox shape must be (N, 8, 3) or (B, N, 8, 3). Got FT)rl   r   r   r+   rn   r6   r;   r   r7   ro   r   rp   r9   rq   rr   rs   rt   r   r   r   rv   V  s    
,zBoxes3D.__init__rw   c                 C   s2   | j dd}|d |d |d   }}}|||fS )a)  Compute boxes heights and widths.

        Returns:
            - Boxes depths, shape of :math:`(N,)` or :math:`(B,N)`.
            - Boxes heights, shape of :math:`(N,)` or :math:`(B,N)`.
            - Boxes widths, shape of :math:`(N,)` or :math:`(B,N)`.

        Example:
            >>> boxes_xyzxyz = torch.tensor([[ 0,  1,  2, 10, 21, 32], [3, 4, 5, 43, 54, 65]])
            >>> boxes3d = Boxes3D.from_tensor(boxes_xyzxyz)
            >>> boxes3d.get_boxes_shape()
            (tensor([30., 60.]), tensor([20., 50.]), tensor([10., 40.]))
        xyzwhd)rO   rT   .r   .   )rz   )ru   Zboxes_xyzwhdr{   r|   Zdepthsr   r   r   r}   o  s    zBoxes3D.get_boxes_shapexyzxyzrN   c                 C   s  d|j   krdkr(n n|jd dks:td|j d|j dk}|rL|n|d}| rb|n| }|d |d	 |d
   }}}| }|dkr|d |d  }|d |d	  }	|d |d
  }
n~|dkr|d |d  d }|d |d	  d }	|d |d
  d }
n6|dkr8|d |d |d   }
}	}ntd| |r|dk rbtd|	dk rxtd|
dk rtdt|||||	|
}|r|n|	d}| |d|dS )a  Helper method to easily create :class:`Boxes3D` from 3D boxes stored in another format.

        Args:
            boxes: 3D boxes, shape of :math:`(N,6)` or :math:`(B,N,6)`.
            mode: The format in which the 3D boxes are provided.

                * 'xyzxyz': boxes are assumed to be in the format ``xmin, ymin, zmin, xmax, ymax, zmax`` where
                  ``width = xmax - xmin``, ``height = ymax - ymin`` and ``depth = zmax - zmin``.
                * 'xyzxyz_plus': similar to 'xyzxyz' mode but where box width, length and depth are defined as
                  ``width = xmax - xmin + 1``, ``height = ymax - ymin + 1`` and ``depth = zmax - zmin + 1``.
                * 'xyzwhd': boxes are assumed to be in the format ``xmin, ymin, zmin, width, height, depth`` where
                  ``width = xmax - xmin``, ``height = ymax - ymin`` and ``depth = zmax - zmin``.

            validate_boxes: check if boxes are valid rectangles or not. Valid rectangles are those with width, height
                and depth >= 1 (>= 2 when mode ends with '_plus' suffix).

        Returns:
            :class:`Boxes3D` class containing the original `boxes` in the format specified by ``mode``.

        Examples:
            >>> boxes_xyzxyz = torch.as_tensor([[0, 3, 6, 1, 4, 8], [5, 1, 3, 8, 4, 9]])
            >>> boxes = Boxes3D.from_tensor(boxes_xyzxyz, mode='xyzxyz')
            >>> boxes.data  # (2, 8, 3)
            tensor([[[0., 3., 6.],
                     [0., 3., 6.],
                     [0., 3., 6.],
                     [0., 3., 6.],
                     [0., 3., 7.],
                     [0., 3., 7.],
                     [0., 3., 7.],
                     [0., 3., 7.]],
            <BLANKLINE>
                    [[5., 1., 3.],
                     [7., 1., 3.],
                     [7., 3., 3.],
                     [5., 3., 3.],
                     [5., 1., 8.],
                     [7., 1., 8.],
                     [7., 3., 8.],
                     [5., 3., 8.]]])
        r   r   r5   r   z,BBox shape must be (N, 6) or (B, N, 6). Got r'   r   rE   rF   rU   r   rT   r   r   r   rB   r   rR   rX   rY   z%Some boxes have negative depths or 0.F)rh   rO   )
r9   r   r;   r:   r6   r7   rZ   r]   rf   r^   )r   r   rO   rP   r_   r>   r?   rb   r@   rA   rc   Zhexahedronsr   r   r   r     s8    +(


zBoxes3D.from_tensor)rO   r   c                 C   sx  | j jrtd| jr| j n
| j d}tj|jdd|jddgdd	|j
d |j
d d}| }|dv rrnl|dv r|d	 |d
  d }|d |d  d }|d |d  d }||d	< ||d< ||d< ntd| |dv rtjg d|j|jd}|| }|dr^|d
 |d |d   }}	}
|d	 |d |d   }}}t||	|
|||}| jrj|n|d}|S )a  Cast :class:`Boxes3D` to a tensor. ``mode`` controls which 3D boxes format should be use to represent
        boxes in the tensor.

        Args:
            mode: The format in which the boxes are provided.

                * 'xyzxyz': boxes are assumed to be in the format ``xmin, ymin, zmin, xmax, ymax, zmax`` where
                  ``width = xmax - xmin``, ``height = ymax - ymin`` and ``depth = zmax - zmin``.
                * 'xyzxyz_plus': similar to 'xyzxyz' mode but where box width, length and depth are defined as
                   ``width = xmax - xmin + 1``, ``height = ymax - ymin + 1`` and ``depth = zmax - zmin + 1``.
                * 'xyzwhd': boxes are assumed to be in the format ``xmin, ymin, zmin, width, height, depth`` where
                  ``width = xmax - xmin``, ``height = ymax - ymin`` and ``depth = zmax - zmin``.
                * 'vertices': boxes are defined by their vertices points in the following ``clockwise`` order:
                  *front-top-left, front-top-right, front-bottom-right, front-bottom-left, back-top-left,
                  back-top-right, back-bottom-right,  back-bottom-left*. Vertices coordinates are in (x,y, z) order.
                  Finally, box width, height and depth are defined as ``width = xmax - xmin``, ``height = ymax - ymin``
                  and ``depth = zmax - zmin``.
                * 'vertices_plus': similar to 'vertices' mode but where box width, length and depth are defined as
                  ``width = xmax - xmin + 1`` and ``height = ymax - ymin + 1``.

        Returns:
            3D Boxes tensor in the ``mode`` format. The shape depends with the ``mode`` value:

                * 'vertices' or 'verticies_plus': :math:`(N, 8, 3)` or :math:`(B, N, 8, 3)`.
                * Any other value: :math:`(N, 6)` or :math:`(B, N, 6)`.

        Note:
            It is currently non-differentiable due to a bug. See github issue
            `#1304 <https://github.com/kornia/kornia/issues/1396>`_.

        Examples:
            >>> boxes_xyzxyz = torch.as_tensor([[0, 3, 6, 1, 4, 8], [5, 1, 3, 8, 4, 9]])
            >>> boxes = Boxes3D.from_tensor(boxes_xyzxyz, mode='xyzxyz')
            >>> assert (boxes.to_tensor(mode='xyzxyz') == boxes_xyzxyz).all()
        a  Boxes3D.to_tensor doesn't support computing gradients since they aren't accurate. Please, create boxes from tensors with `requires_grad=False`. This is a known bug. Help is needed to fix it. For more information, see https://github.com/kornia/kornia/issues/1396.r   r   rd   rB   r   )r   r   )r   rQ   rS   rT   rE   r   rF   r   rU   rR   )r   rQ   )r   r   r   rB   rB   rB   rC   rQ   )rr   r   r   rq   r:   r   r   r   r   r8   r   rZ   r;   r   rD   r   r[   rf   r^   )ru   rO   r   r   r@   rA   rc   r   r>   r?   rb   r   r   r   rz     s8    $"

zBoxes3D.to_tensor)rc   rA   r@   r   c                 C   s0  | j jrtd| jrJtj| j jd | j jd |||f| j j| j jd}n(tj| j jd |||f| j j| j jd}| 	d}|ddddf 
d| |ddddf 
d| |dd	ddf 
d| t|d
||||d
d  D ]:\}}d||d	 |d |d |d |d |d f< q|S )u  Convert ·D boxes to masks. Covered area is 1 and the remaining is 0.

        Args:
            depth: depth of the masked image/images.
            height: height of the masked image/images.
            width: width of the masked image/images.

        Returns:
            the output mask tensor, shape of :math:`(N, depth, width, height)` or :math:`(B,N, depth, width, height)`
             and dtype of :func:`Boxes3D.dtype` (it can be any floating point dtype).

        Note:
            It is currently non-differentiable.

        Examples:
            >>> boxes = Boxes3D(torch.tensor([[  # Equivalent to boxes = Boxes.3Dfrom_tensor([[1,1,1,3,3,2]])
            ...     [1., 1., 1.],
            ...     [3., 1., 1.],
            ...     [3., 3., 1.],
            ...     [1., 3., 1.],
            ...     [1., 1., 2.],
            ...     [3., 1., 2.],
            ...     [3., 3., 2.],
            ...     [1., 3., 2.],
            ... ]]))  # 1x8x3
            >>> boxes.to_mask(4, 5, 5)
            tensor([[[[0., 0., 0., 0., 0.],
                      [0., 0., 0., 0., 0.],
                      [0., 0., 0., 0., 0.],
                      [0., 0., 0., 0., 0.],
                      [0., 0., 0., 0., 0.]],
            <BLANKLINE>
                     [[0., 0., 0., 0., 0.],
                      [0., 1., 1., 1., 0.],
                      [0., 1., 1., 1., 0.],
                      [0., 1., 1., 1., 0.],
                      [0., 0., 0., 0., 0.]],
            <BLANKLINE>
                     [[0., 0., 0., 0., 0.],
                      [0., 1., 1., 1., 0.],
                      [0., 1., 1., 1., 0.],
                      [0., 1., 1., 1., 0.],
                      [0., 0., 0., 0., 0.]],
            <BLANKLINE>
                     [[0., 0., 0., 0., 0.],
                      [0., 0., 0., 0., 0.],
                      [0., 0., 0., 0., 0.],
                      [0., 0., 0., 0., 0.],
                      [0., 0., 0., 0., 0.]]]])
        r   r   rB   r   r   .Nr   r   r5   r   r   r   )rr   r   r   rq   r   rK   r   r   rD   rz   r   r   r8   r   r   )ru   rc   rA   r@   r   Zclipped_boxes_xyzxyzr   Z
box_xyzxyzr   r   r   r     s2    3
 ,zBoxes3D.to_maskFr   c                 C   sd   d|j   krdkr,n n|jdd dkr>td|j dt| j|}|rX|| _| S t|dd	S )
a  Apply a transformation matrix to the 3D boxes.

        Args:
            M: The transformation matrix to be applied, shape of :math:`(4, 4)` or :math:`(B, 4, 4)`.
            inplace: do transform in-place and return self.

        Returns:
            The transformed boxes.
        r   r   r   N)r   r   zAThe transformation matrix shape must be (4, 4) or (B, 4, 4). Got r'   Fr   )r9   r   r;   r=   rr   r   r   r   r   r   r   r  s    
,zBoxes3D.transform_boxesr   c                 C   s   | j |ddS )z2Inplace version of :func:`Boxes3D.transform_boxes`Tr   r   r   r   r   r   r     s    zBoxes3D.transform_boxes_c                 C   s   | j S r   r   r   r   r   r   r     s    zBoxes3D.datac                 C   s   | j S r   r   r   r   r   r   rO     s    zBoxes3D.modec                 C   s   | j jS r   r   r   r   r   r   rD     s    zBoxes3D.devicec                 C   s   | j jS r   r   r   r   r   r   r     s    zBoxes3D.dtyper   c                 C   s.   |durt |std| jj||d| _| S r   r   r   r   r   r   r     s    z
Boxes3D.to)Tr   )r   T)r   )F)NN)r   r   r   r   r   r   r   r   rv   r   r}   r   r   rz   r   r   r   r   r   r   rO   rD   r   r   r   r   r   r   r   r   E  s,    NMU)r   )rM   T)typingr   r   r   r   r   r   Zkornia.corer   Zkornia.geometry.bboxr   Zkornia.geometry.linalgr	   Zkornia.utilsr
   __all__r   r   r   r   r   r2   r=   rL   ra   rf   r   Zjitscriptr   r   r   r   r   <module>   s6   (;!   !