
    iD                        d dl Z d dlZd dlmZmZmZmZ d dlZd dlmZm	Z	 d dl
mZ d dlmZ dddd	Zesd
dgZd7dedede	j$                  j&                  j(                  fdZ G d dej                  j,                        Z G d dej                  j,                        Z G d dej                  j,                        Zd8dededefdZd9dedeedf   defdZd:dededefdZd;ded edefd!Z  G d" d#e	j,                        Z! G d$ d%e	j,                        Z" G d& d'e	j,                        Z# G d( d)e#      Z$d*ed+edefd,Z%d-ed.ede	j,                  d+edef
d/Z&d<d0ed1eed2      defd3Z'	 	 	 d=d-ed.ed4ed5   d1eed2      d+edefd6Z(y)>    N)List
NamedTupleOptionalUnion)Tensornn)Literal)_TORCHVISION_AVAILABLESqueezeNet1_1_WeightsAlexNet_WeightsVGG16_Weights)squeezenet1_1alexnetvgg16)learned_perceptual_image_patch_similarity_get_tv_model_featuresnet
pretrainedreturnc                 "   t         st        d      ddl}|rPt        |j                  t
        |          } t        |j                  |       |j                        }|j                  S  t        |j                  |       d      }|j                  S )aA  Get torchvision network.

    Args:
        net: Name of network
        pretrained: If pretrained weights should be used

    >>> _ = _get_tv_model_features("alexnet", pretrained=True)
    >>> _ = _get_tv_model_features("squeezenet1_1", pretrained=True)
    >>> _ = _get_tv_model_features("vgg16", pretrained=True)

    zSTorchvision is not installed. Please install torchvision to use this functionality.r   N)weights)r
   ModuleNotFoundErrortorchvisiongetattrmodels_weight_mapDEFAULTfeatures)r   r   r   model_weightsmodels        x/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/torchmetrics/functional/image/lpips.pyr   r   -   s     "!"wxx 2 2K4DE0**C09N9NO >> 1**C0>>>    c                   B     e Zd ZdZd	dededdf fdZdedefdZ xZ	S )

SqueezeNetzSqueezeNet implementation.requires_gradr   r   Nc           
         t         
|           t        d|      }d| _        g }t	        d      t	        dd      t	        dd      t	        dd      t	        dd      t	        dd      t	        dd	      g}|D ]V  }t
        j                  j                         }|D ]   }|j                  t        |      ||          " |j                  |       X t        j                  |      | _        |s| j                         D ]	  }	d
|	_         y y )Nr               
            F)super__init__r   N_slicesrangetorchr   
Sequential
add_modulestrappend
ModuleListslices
parametersr%   )selfr%   r   pretrained_featuresr9   feature_rangesfeature_rangeseqiparam	__class__s             r!   r0   zSqueezeNet.__init__H   s    4_jQ(E!QKq!eArlERTVXM[`aceg[hjoprtvjwx+M((%%'C"s1v':1'=> #MM#	 , mmF+*&+# + r"   xc                      G d dt               }g }| j                  D ]  } ||      }|j                  |         || S )Process input.c                   T    e Zd ZU eed<   eed<   eed<   eed<   eed<   eed<   eed<   y)	*SqueezeNet.forward.<locals>._SqueezeOutputrelu1relu2relu3relu4relu5relu6relu7N__name__
__module____qualname__r   __annotations__ r"   r!   _SqueezeOutputrG   ]   s%    MMMMMMMr"   rU   )r   r9   r7   )r;   rC   rU   relusslice_s        r!   forwardzSqueezeNet.forwardZ   sD    	Z 	 kkFq	ALLO " u%%r"   FT
rP   rQ   rR   __doc__boolr0   r   r   rX   __classcell__rB   s   @r!   r$   r$   E   s4    $,d , ,PT ,$& &J &r"   r$   c                   B     e Zd ZdZd	dededdf fdZdedefdZ xZ	S )
AlexnetzAlexnet implementation.r%   r   r   Nc                    t         |           t        d|      }t        j                  j                         | _        t        j                  j                         | _        t        j                  j                         | _        t        j                  j                         | _	        t        j                  j                         | _
        d| _        t        d      D ]*  }| j                  j                  t        |      ||          , t        dd      D ]*  }| j                  j                  t        |      ||          , t        dd      D ]*  }| j                  j                  t        |      ||          , t        dd      D ]*  }| j                  j                  t        |      ||          , t        dd      D ]*  }| j                  j                  t        |      ||          , |s| j                         D ]	  }d|_         y y )Nr   r)   r(   r*   r+   r-   Fr/   r0   r   r3   r   r4   slice1slice2slice3slice4slice5r1   r2   r5   r6   r:   r%   )r;   r%   r   alexnet_pretrained_featuresrC   rA   rB   s         r!   r0   zAlexnet.__init__p   s   &<Y
&S#hh))+hh))+hh))+hh))+hh))+qAKK""3q6+Fq+IJ q!AKK""3q6+Fq+IJ q!AKK""3q6+Fq+IJ q"AKK""3q6+Fq+IJ r2AKK""3q6+Fq+IJ *&+# + r"   rC   c                     | j                  |      }|}| j                  |      }|}| j                  |      }|}| j                  |      }|}| j	                  |      }|} G d dt
              } ||||||      S )rE   c                   @    e Zd ZU eed<   eed<   eed<   eed<   eed<   y)(Alexnet.forward.<locals>._AlexnetOutputsrH   rI   rJ   rK   rL   NrO   rT   r"   r!   _AlexnetOutputsrk      s    MMMMMr"   rl   rc   rd   re   rf   rg   r   )	r;   rC   hh_relu1h_relu2h_relu3h_relu4h_relu5rl   s	            r!   rX   zAlexnet.forward   s}    KKNKKNKKNKKNKKN	j 	 w'7KKr"   rY   rZ   r^   s   @r!   r`   r`   m   s7    !,d , ,PT ,0L LJ Lr"   r`   c                   B     e Zd ZdZd	dededdf fdZdedefdZ xZ	S )
Vgg16zVgg16 implementation.r%   r   r   Nc                    t         |           t        d|      }t        j                  j                         | _        t        j                  j                         | _        t        j                  j                         | _        t        j                  j                         | _	        t        j                  j                         | _
        d| _        t        d      D ]*  }| j                  j                  t        |      ||          , t        dd      D ]*  }| j                  j                  t        |      ||          , t        dd      D ]*  }| j                  j                  t        |      ||          , t        dd      D ]*  }| j                  j                  t        |      ||          , t        dd      D ]*  }| j                  j                  t        |      ||          , |s| j                         D ]	  }d|_         y y )	Nr   r)      	            Frb   )r;   r%   r   vgg_pretrained_featuresrC   rA   rB   s         r!   r0   zVgg16.__init__   s   "8*"Mhh))+hh))+hh))+hh))+hh))+qAKK""3q6+B1+EF q!AKK""3q6+B1+EF q"AKK""3q6+B1+EF r2AKK""3q6+B1+EF r2AKK""3q6+B1+EF *&+# + r"   rC   c                     | j                  |      }|}| j                  |      }|}| j                  |      }|}| j                  |      }|}| j	                  |      }|} G d dt
              } ||||||      S )rE   c                   @    e Zd ZU eed<   eed<   eed<   eed<   eed<   y)"Vgg16.forward.<locals>._VGGOutputsrelu1_2relu2_2relu3_3relu4_3relu5_3NrO   rT   r"   r!   _VGGOutputsr      s    OOOOOr"   r   rm   )	r;   rC   rn   	h_relu1_2	h_relu2_2	h_relu3_3	h_relu4_3	h_relu5_3r   s	            r!   rX   zVgg16.forward   s}    KKN	KKN	KKN	KKN	KKN		* 	 9iIyQQr"   rY   rZ   r^   s   @r!   ru   ru      s7    ,d , ,PT ,0R RJ Rr"   ru   in_tenskeep_dimc                 ,    | j                  ddg|      S )z2Spatial averaging over height and width of images.r(      )keepdimmean)r   r   s     r!   _spatial_averager      s    <<A<11r"   out_hw.c                 >     t        j                  |dd      |       S )z+Upsample input with bilinear interpolation.bilinearF)sizemodealign_corners)r   Upsample)r   r   s     r!   	_upsampler      s    I2;;F5I'RRr"   in_featepsc                 n    t        j                  |t        j                  | dz  dd      z         }| |z  S )zNormalize input tensor.r(      T)dimr   )r3   sqrtsum)r   r   norm_factors      r!   _normalize_tensorr      s1    **S599WaZQ#MMNK[  r"   rC   r   c                    | j                   d   |kD  r@| j                   d   |kD  r.t        j                  j                  j	                  | ||fd      S t        j                  j                  j	                  | ||fdd      S )zlhttps://github.com/toshas/torch-fidelity/blob/master/torch_fidelity/sample_similarity_lpips.py#L127C22-L132.area)r   r   F)r   r   )shaper3   r   
functionalinterpolate)rC   r   s     r!   _resize_tensorr      sp    wwr{TaggbkD0xx""..q4,V.LL88**1tTl[`*aar"   c                   J     e Zd ZU dZeed<   eed<   d fdZdedefdZ xZS )	ScalingLayerzScaling layer.shiftscaler   c                     t         |           | j                  dt        j                  g d      d d d d d f   d       | j                  dt        j                  g d      d d d d d f   d       y )Nr   )gQgI+gMbȿF)
persistentr   )gZd;O?gy&1?g?)r/   r0   register_bufferr3   r   )r;   rB   s    r!   r0   zScalingLayer.__init__   sp    Well3K&LTSTVZ\`M`&anstWell3H&I$PQSWY]J]&^kpqr"   inpc                 :    || j                   z
  | j                  z  S rE   )r   r   )r;   r   s     r!   rX   zScalingLayer.forward   s    djj DJJ..r"   )r   N)	rP   rQ   rR   r[   r   rS   r0   rX   r]   r^   s   @r!   r   r      s)    MMr
/6 /f /r"   r   c            	       F     e Zd ZdZd
dedededdf fdZdedefd	Z xZ	S )NetLinLayerz,A single linear layer which does a 1x1 conv.chn_inchn_outuse_dropoutr   Nc           	          t         |           |rt        j                         gng }|t        j                  ||dddd      gz  }t        j
                  | | _        y )Nr   r   F)stridepaddingbias)r/   r0   r   DropoutConv2dr4   r    )r;   r   r   r   layersrB   s        r!   r0   zNetLinLayer.__init__   sV    #."**,BIIfgqAEJ
 	
 ]]F+
r"   rC   c                 $    | j                  |      S r   )r    )r;   rC   s     r!   rX   zNetLinLayer.forward  s    zz!}r"   )r   F)
rP   rQ   rR   r[   intr\   r0   r   rX   r]   r^   s   @r!   r   r      s;    6,s ,S ,4 ,TX , F r"   r   c                        e Zd Z	 	 	 	 	 	 	 	 	 ddeded   dedededed	ee   d
edee   ddf fdZ	 dde	de	dedede
e	ee	ee	   f   f   f
dZ xZS )_LPIPSNr   r   alexvggsqueezespatial	pnet_rand	pnet_tuner   
model_path	eval_moderesizer   c
           	      .   t         |           || _        || _        || _        || _        |	| _        t               | _        | j                  dv rt        }
g d| _
        n=| j                  dk(  rt        }
g d| _
        n| j                  dk(  rt        }
g d| _
        t        | j                        | _         
| j                   | j                        | _        t!        | j                  d   |	      | _        t!        | j                  d
   |	      | _        t!        | j                  d   |	      | _        t!        | j                  d   |	      | _        t!        | j                  d   |	      | _        | j"                  | j$                  | j&                  | j(                  | j*                  g| _        | j                  dk(  rit!        | j                  d   |	      | _        t!        | j                  d   |	      | _        | xj,                  | j.                  | j0                  gz  c_        t3        j4                  | j,                        | _        |r|_t6        j8                  j;                  t6        j8                  j=                  t?        j@                  | j                        dd| d            }| jC                  tE        jF                  |d      d       |r| jI                          | j                  s| jK                         D ]	  }d|_&         yy)a  Initializes a perceptual loss torch.nn.Module.

        Args:
            pretrained: This flag controls the linear layers should be pretrained version or random
            net: Indicate backbone to use, choose between ['alex','vgg','squeeze']
            spatial: If input should be spatial averaged
            pnet_rand: If backbone should be random or use imagenet pre-trained weights
            pnet_tune: If backprop should be enabled for both backbone and linear layers
            use_dropout: If dropout layers should be added
            model_path: Model path to load pretained models from
            eval_mode: If network should be in evaluation mode
            resize: If input should be resized to this size

        )r   r   )@            r   r   )r        r   r   r   )r   r   r   r   r   r   r   )r   r%   r   )r   r   r(   r   rw   r)      Nz..zlpips_models/z.pthcpu)map_locationF)strict)'r/   r0   	pnet_typer   r   r   r   r   scaling_layerru   chnsr`   r$   lenLr   r   lin0lin1lin2lin3lin4linslin5lin6r   r8   ospathabspathjoininspectgetfileload_state_dictr3   loadevalr:   r%   )r;   r   r   r   r   r   r   r   r   r   net_typerA   rB   s               r!   r0   z_LPIPS.__init__
  sC   4 	"")^>>--H0DI^^v%H0DI^^y(!H:DITYY4>>'9X		!+F			!+F			!+F			!+F			!+F	YY		499diiK	>>Y&#DIIaLkJDI#DIIaLkJDIII$))TYY//IMM$)),	!WW__GGLL!?WZV[[_G`a
   JU!KTY ZIIK~~*&+# + r"   in0in1retperlayer	normalizec                 ^   |rd|z  dz
  }d|z  dz
  }| j                  |      | j                  |      }}| j                  .t        || j                        }t        || j                        }| j                  j	                  |      | j                  j	                  |      }}i i i }}
}	t        | j                        D ]6  }t        ||         t        ||         c|	|<   |
|<   |	|   |
|   z
  dz  ||<   8 g }t        | j                        D ]  }| j                  rI|j                  t         | j                  |   ||         t        |j                  dd                     X|j                  t         | j                  |   ||         d              t        |      }|r||fS |S )Nr(   r   )r   )r   T)r   )r   r   r   r   rX   r2   r   r   r   r7   r   r   tupler   r   r   )r;   r   r   r   r   	in0_input	in1_inputouts0outs1feats0feats1diffskkresvals                  r!   rX   z_LPIPS.forwardU  s    c'A+Cc'A+C  $11#68J8J38O9	 ;;"&yt{{CI&yt{{CIxx''	2DHH4D4DY4Ou "B-B%6uRy%ACTUZ[]U^C_"F2Jr
fRj0Q6E"I   -B||

9]TYYr]59%=eCIIVWVXMFZ[\

+MDIIbM%),DtTU	   #h:
r"   )	Tr   FFFTNTN)FF)rP   rQ   rR   r\   r	   r   r6   r   r0   r   r   r   r   rX   r]   r^   s   @r!   r   r   	  s      17 $( $I,I, -.I, 	I,
 I, I, I, SMI, I, I, 
I,X V[   & 59 NR 	vuVT&\122	3 r"   r   c                   ,     e Zd ZdZdedd f fdZ xZS )_NoTrainLpipsz8Wrapper to make sure LPIPS never leaves evaluation mode.r   r   c                 "    t         |   d      S )z.Force network to always be in evaluation mode.F)r/   train)r;   r   rB   s     r!   r  z_NoTrainLpips.train{  s    w}U##r"   )rP   rQ   rR   r[   r\   r  r]   r^   s   @r!   r  r  x  s    B$$ $? $ $r"   r  imgr   c                     |r(| j                         dk  xr& | j                         dk\  n| j                         dk\  }| j                  dk(  xr | j                  d   dk(  xr |S )z1Check that input is a valid image to the network.g      ?g        r   rw   r   r   )maxminndimr   )r  r   value_checks      r!   
_valid_imgr    sV    ;D#'')s"7swwyC'7#'')WY/K88q=>SYYq\Q.>;>r"   img1img2c                 J   t        | |      rt        ||      sst        d| j                   d|j                   d| j                         | j	                         g d|j                         |j	                         g d|rddgnddg d       || ||	      j                         S )
NzeExpected both input arguments to be normalized tensors with shape [N, 3, H, W]. Got input with shape z and z and values in range z+ when all values are expected to be in the r   r   r   z range.)r   )r  
ValueErrorr   r	  r  r   )r  r  r   r   s       r!   _lpips_updater    s    tY'JtY,G%%)ZZLdjj\ BTXXZ()
DHHJ/G.H I&09q!fAw%GwP
 	
 tTY/7799r"   scores	reduction)r   r   nonec                     |dk(  r| j                         S |dk(  r| j                         S |dk(  s|| S t        d|       )Nr   r   r  zInvalid reduction type: )r   r   r  )r  r  s     r!   _lpips_computer    sO    F{{}Ezz|Fi/
/	{;
<<r"   r   r   c                     t        |      j                  | j                  | j                        }t	        | |||      }t        ||      S )a[  The Learned Perceptual Image Patch Similarity (`LPIPS_`) calculates perceptual similarity between two images.

    LPIPS essentially computes the similarity between the activations of two image patches for some pre-defined network.
    This measure has been shown to match human perception well. A low LPIPS score means that image patches are
    perceptual similar.

    Both input image patches are expected to have shape ``(N, 3, H, W)``. The minimum size of `H, W` depends on the
    chosen backbone (see `net_type` arg).

    Args:
        img1: first set of images
        img2: second set of images
        net_type: str indicating backbone network type to use. Choose between `'alex'`, `'vgg'` or `'squeeze'`
        reduction: str indicating how to reduce over the batch dimension. Choose between `'sum'`, `'mean'`, `'none'`
            or `None`.
        normalize: by default this is ``False`` meaning that the input is expected to be in the [-1,1] range. If set
            to ``True`` will instead expect input to be in the ``[0,1]`` range.

    Example:
        >>> from torch import rand
        >>> from torchmetrics.functional.image.lpips import learned_perceptual_image_patch_similarity
        >>> img1 = (rand(10, 3, 100, 100) * 2) - 1
        >>> img2 = (rand(10, 3, 100, 100) * 2) - 1
        >>> learned_perceptual_image_patch_similarity(img1, img2, net_type='squeeze')
        tensor(0.1005)

        >>> from torch import rand, Generator
        >>> from torchmetrics.functional.image.lpips import learned_perceptual_image_patch_similarity
        >>> gen = Generator().manual_seed(42)
        >>> img1 = (rand(2, 3, 100, 100, generator=gen) * 2) - 1
        >>> img2 = (rand(2, 3, 100, 100, generator=gen) * 2) - 1
        >>> learned_perceptual_image_patch_similarity(img1, img2, net_type='squeeze', reduction='none')
        tensor([0.1024, 0.0938])

    )r   )devicedtype)r  tor  r  r  r  )r  r  r   r  r   r   losss          r!   r   r     sD    T H
%
(
(4::
(
NCtS)4D$	**r"   )F)T))r   r   )g:0yE>)r   r   )r   r   F))r   r   typingr   r   r   r   r3   r   r   typing_extensionsr	   torchmetrics.utilities.importsr
   r   __doctest_skip__r6   r\   modules	containerr4   r   Moduler$   r`   ru   r   r   r   r   floatr   r   r   r   r   r  r  r  r  r   rT   r"   r!   <module>r$     s,  2  	 4 4   % A -  CE]^  "**BVBVBaBa 0%& %&P/Lehhoo /Ld/REHHOO /Rd2f 2 2 2
Sv SuS#X Sf S
!v !E !V !bf bC b b/299 / ")) "lRYY l^$F $?F ?t ? ?: :f :299 : :RX :=6 =hw?T7U.V =dj = 39:@,+
,+
,+ ./,+  567	,+
 ,+ ,+r"   