
    ^jX                         d dl mZ d dlmZ d dlZ G d dee      Zeeef   ZdefdZdefdZ	d	ej                  defd
Zd	ej                  defdZy)    )Enum)UnionNc                       e Zd ZdZdZdZdZy)FormatNCHWNHWCNCLNLCN)__name__
__module____qualname__r   r   r	   r
        ]/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/timm/layers/format.pyr   r      s    DD
C
Cr   r   fmtc                     t        |       } | t         j                  u rd}|S | t         j                  u rd}|S | t         j                  u rd}|S d}|S )zReturn spatial dimension indices for a given tensor format.

    Args:
        fmt: Tensor format (NCHW, NHWC, NCL, or NLC).

    Returns:
        Tuple of spatial dimension indices.
    )   )   )r   r   )r      )r   r
   r	   r   r   dims     r   get_spatial_dimr      se     +C
fjj J 


	
 J	 
	 J Jr   c                 x    t        |       } | t         j                  u rd}|S | t         j                  u rd}|S d}|S )zReturn channel dimension index for a given tensor format.

    Args:
        fmt: Tensor format (NCHW, NHWC, NCL, or NLC).

    Returns:
        Channel dimension index.
    r   r   r   )r   r   r
   r   s     r   get_channel_dimr   &   sK     +C
fkk
 J	 


	 J Jr   xc                    |t         j                  k(  r| j                  dddd      } | S |t         j                  k(  r#| j	                  d      j                  dd      } | S |t         j                  k(  r| j	                  d      } | S )zConvert tensor from NCHW format to specified format.

    Args:
        x: Input tensor in NCHW format.
        fmt: Target format.

    Returns:
        Tensor in target format.
    r   r   r   r   )r   r   permuter
   flatten	transposer	   r   r   s     r   nchw_tor!   9   sz     fkkIIaAq!
 H	 


	IIaL""1a( H 


	IIaLHr   c                    |t         j                  k(  r| j                  dddd      } | S |t         j                  k(  r| j	                  dd      } | S |t         j
                  k(  r"| j	                  dd      j                  dd      } | S )zConvert tensor from NHWC format to specified format.

    Args:
        x: Input tensor in NHWC format.
        fmt: Target format.

    Returns:
        Tensor in target format.
    r   r   r   r   )r   r   r   r
   r   r	   r   r    s     r   nhwc_tor#   L   s~     fkkIIaAq!
 H	 


	IIaO H 


	IIaO%%a+Hr   )enumr   typingr   torchstrr   FormatTr   r   Tensorr!   r#   r   r   r   <module>r*      sr      S$  V
 * &u|| & &u|| & r   