ó
    !Eñia  ã                  óà   • S SK Jr  S SKrS SKJr  S SKr/ SQr\" SSS9r\" SS	S9r S       SS
 jjr	 S                 SS jjr
      SS jr                SS jrg)é    )ÚannotationsN)ÚTypeVar)Úfuse_conv_bn_evalÚfuse_conv_bn_weightsÚfuse_linear_bn_evalÚfuse_linear_bn_weightsÚConvTztorch.nn.modules.conv._ConvNd)ÚboundÚLinearTztorch.nn.Linearc           
     ó   • U R                   (       d  UR                   (       a  [        S5      e[        R                  " U 5      nUR                  b  UR
                  c  [        S5      e[        UR                  UR                  UR                  UR
                  UR                  UR                  UR                  U5      u  Ul        Ul        U$ )a  Fuse a convolutional module and a BatchNorm module into a single, new convolutional module.

Args:
    conv (torch.nn.modules.conv._ConvNd): A convolutional module.
    bn (torch.nn.modules.batchnorm._BatchNorm): A BatchNorm module.
    transpose (bool, optional): If True, transpose the convolutional weight. Defaults to False.

Returns:
    torch.nn.modules.conv._ConvNd: The fused convolutional module.

.. note::
    Both ``conv`` and ``bn`` must be in eval mode, and ``bn`` must have its running buffers computed.
úFusion only for eval!ú3bn.running_mean and bn.running_var must not be None)
ÚtrainingÚAssertionErrorÚcopyÚdeepcopyÚrunning_meanÚrunning_varr   ÚweightÚbiasÚeps)ÚconvÚbnÚ	transposeÚ
fused_convs       ÚR/home/mande/repo/quber/.venv/lib/python3.13/site-packages/torch/nn/utils/fusion.pyr   r      sœ   € ð$ ‡}‡}˜ŸŸÜÐ4Ó5Ð5Ü—’˜tÓ$€Jà	‡�Ñ "§.¡.Ñ"8ÜÐRÓSÐSÜ)=Ø×ÑØ�‰Ø
�‰Ø
�‰Ø
�‰Ø
�	‰	Ø
�‰Øó	*Ñ&€JÔ�z”ð Ðó    c                ó´  • U R                   nUb  UR                   OUn	Uc  [        R                  " U5      nUc  [        R                  " U5      nUc  [        R                  " U5      n[        R                  " X4-   5      n
U(       a"  SS/S/[        U R                  5      S-
  -  -   nO!SS/S/[        U R                  5      S-
  -  -   nXU
-  R                  U5      -  R                  US9nX-
  U
-  U-  U-   R                  U	S9n[        R                  R                  XÀR                  5      [        R                  R                  XÑR                  5      4$ )a�  Fuse convolutional module parameters and BatchNorm module parameters into new convolutional module parameters.

Args:
    conv_w (torch.Tensor): Convolutional weight.
    conv_b (Optional[torch.Tensor]): Convolutional bias.
    bn_rm (torch.Tensor): BatchNorm running mean.
    bn_rv (torch.Tensor): BatchNorm running variance.
    bn_eps (float): BatchNorm epsilon.
    bn_w (Optional[torch.Tensor]): BatchNorm weight.
    bn_b (Optional[torch.Tensor]): BatchNorm bias.
    transpose (bool, optional): If True, transpose the conv weight. Defaults to False.

Returns:
    Tuple[torch.nn.Parameter, torch.nn.Parameter]: Fused convolutional weight and bias.
é   éÿÿÿÿé   ©Údtype)r#   ÚtorchÚ
zeros_likeÚ	ones_likeÚrsqrtÚlenÚshapeÚreshapeÚtoÚnnÚ	ParameterÚrequires_grad)Úconv_wÚconv_bÚbn_rmÚbn_rvÚbn_epsÚbn_wÚbn_br   Úconv_weight_dtypeÚconv_bias_dtypeÚbn_var_rsqrtr)   Úfused_conv_wÚfused_conv_bs                 r   r   r   :   sI  € ð2 Ÿ™ÐØ&,Ñ&8�f—l’lÐ>O€OØ�~Ü×!Ò! %Ó(ˆØ�|Ü�Š˜uÓ%ˆØ�|Ü×Ò Ó&ˆÜ—;’;˜u™~Ó.€LæØ�B�˜1˜#¤ V§\¡\Ó!2°QÑ!6Ñ7Ñ7‰à�Q�˜1˜#¤ V§\¡\Ó!2°QÑ!6Ñ7Ñ7ˆà \Ñ1×:Ñ:¸5ÓAÑA×EÑEØð Fð €Lð ‘^ |Ñ3°dÑ:¸TÑA×EÑEØð Fð €Lô
 	�‰×Ñ˜<×)=Ñ)=Ó>Ü�‰×Ñ˜<×)=Ñ)=Ó>ðð r   c           	     ó>  • U R                   (       d  UR                   (       a  [        S5      e[        R                  " U 5      n U R                  UR
                  :w  a5  UR
                  S:w  a%  [        SU R                   SUR
                   35      eUR                  b  UR                  c  [        S5      e[        UR                  UR                  UR                  UR                  UR                  UR                  UR                  5      u  Ul	        Ul
        U$ )as  Fuse a linear module and a BatchNorm module into a single, new linear module.

Args:
    linear (torch.nn.Linear): A Linear module.
    bn (torch.nn.modules.batchnorm._BatchNorm): A BatchNorm module.

Returns:
    torch.nn.Linear: The fused linear module.

.. note::
    Both ``linear`` and ``bn`` must be in eval mode, and ``bn`` must have its running buffers computed.
r   r   zaTo fuse, linear.out_features == bn.num_features or bn.num_features == 1, got linear.out_features=z and bn.num_features=r   )r   r   r   r   Úout_featuresÚnum_featuresr   r   r   r   r   r   )Úlinearr   Úfused_linears      r   r   r   o   sõ   € ð  ‡‡˜"Ÿ+Ÿ+ÜÐ4Ó5Ð5Ü—=’= Ó(€Lð	ð ×Ñ˜bŸo™oÓ-°"·/±/ÀQÓ2FÜð'Ø'-×':Ñ':Ð&;Ð;PÐQS×Q`ÑQ`ÐPaðcó
ð 	
ð
 
‡�Ñ "§.¡.Ñ"8ÜÐRÓSÐSÜ-CØ×ÑØ×ÑØ
�‰Ø
�‰Ø
�‰Ø
�	‰	Ø
�‰ó.Ñ*€LÔ˜Ô*ð Ðr   c                ó¶  • U R                   nUb  UR                   OUnUc  [        R                  " U5      nU[        R                  " X4-   5      -  n	X	R	                  S5      R                  US9-  n
X-
  U	-  U-   R                  US9n[        R                  R                  X R                  5      [        R                  R                  X±R                  5      4$ )a  Fuse linear module parameters and BatchNorm module parameters into new linear module parameters.

Args:
    linear_w (torch.Tensor): Linear weight.
    linear_b (Optional[torch.Tensor]): Linear bias.
    bn_rm (torch.Tensor): BatchNorm running mean.
    bn_rv (torch.Tensor): BatchNorm running variance.
    bn_eps (float): BatchNorm epsilon.
    bn_w (torch.Tensor): BatchNorm weight.
    bn_b (torch.Tensor): BatchNorm bias.

Returns:
    Tuple[torch.nn.Parameter, torch.nn.Parameter]: Fused linear weight and bias.
r    r"   )	r#   r$   r%   r'   Ú	unsqueezer+   r,   r-   r.   )Úlinear_wÚlinear_br1   r2   r3   r4   r5   Úlinear_weight_dtypeÚlinear_bias_dtypeÚbn_scaleÚfused_wÚfused_bs               r   r   r   ¢   sË   € ð. #Ÿ.™.ÐØ*2Ñ*>˜ŸšÐDWÐØÑÜ×#Ò# EÓ*ˆØ”e—k’k %¡.Ó1Ñ1€Hà×+Ñ+¨BÓ/×2Ñ2Ð9LÐ2ÐMÑM€GØÑ  HÑ,¨tÑ3×7Ñ7Ð>OÐ7ÐP€Gä�8‰8×Ñ˜g×'=Ñ'=Ó>ÄÇÁ×@RÑ@RØ×'Ñ'óAð ð r   )F)r   r	   r   ú%torch.nn.modules.batchnorm._BatchNormr   ÚboolÚreturnr	   )r/   útorch.Tensorr0   útorch.Tensor | Noner1   rL   r2   rL   r3   Úfloatr4   rM   r5   rM   r   rJ   rK   ú-tuple[torch.nn.Parameter, torch.nn.Parameter])r>   r   r   rI   rK   r   )rB   rL   rC   rM   r1   rL   r2   rL   r3   rN   r4   rL   r5   rL   rK   rO   )Ú
__future__r   r   Útypingr   r$   Ú__all__r	   r   r   r   r   r   © r   r   Ú<module>rT      s5  ðÝ "ã Ý ã ò€ñ 	�Ð>Ñ?€Ù
�)Ð#4Ñ
5€ð ð#Ø
ð#à-ð#ð ð#ð õ	#ð\ ð2Øð2àð2ð ð2ð ð	2ð
 ð2ð ð2ð ð2ð ð2ð 3õ2ðj0Øð0à-ð0ð ô0ðf"Øð"à!ð"ð ð"ð ð	"ð
 ð"ð ð"ð ð"ð 3õ"r   