ó
    !Eñi»™  ã                   óX  • S SK Jr  S SKrS SKJr  S SKJrJr  S SKJ	r	J
r
Jr  SSKJr  SSKJr  SS	KJr  / S
Qr " S S\5      r " S S\5      r " S S\\5      r " S S\5      r " S S\\5      r " S S\5      r " S S\\5      r " S S\5      r " S S\\5      r " S S\5      rg)é    )ÚAnyN)ÚTensor)Ú
functionalÚinit)Ú	ParameterÚUninitializedBufferÚUninitializedParameteré   )ÚSyncBatchNorm)ÚLazyModuleMixin)ÚModule)ÚBatchNorm1dÚLazyBatchNorm1dÚBatchNorm2dÚLazyBatchNorm2dÚBatchNorm3dÚLazyBatchNorm3dr   c                   óØ   ^ • \ rS rSr% SrSr/ SQr\\S'   \	\S'   \	S-  \S'   \
\S	'   \
\S
'         SS\S\	S\	S-  S	\
S
\
SS4U 4S jjjrSS jrSS jrS rS r  SU 4S jjrSrU =r$ )Ú	_NormBaseé   z,Common base of _InstanceNorm and _BatchNorm.é   )Útrack_running_statsÚmomentumÚepsÚnum_featuresÚaffiner   r   Nr   r   r   Úreturnc                 óŽ  >• XgS.n[         TU ]  5         Xl        X l        X0l        X@l        XPl        U R
                  (       aK  [        [        R                  " U40 UD65      U l
        [        [        R                  " U40 UD65      U l        O$U R                  SS 5        U R                  SS 5        U R                  (       a·  U R                  S[        R                  " U40 UD65        U R                  S[        R                  " U40 UD65        U   U   U R                  S[        R                   "  SS[        R"                  0UR%                  5        V	V
s0 s H  u  pšU	S:w  d  M  Xš_M     sn
n	D65        U   O6U R                  SS 5        U R                  SS 5        U R                  SS 5        U R'                  5         g s  sn
n	f )	N©ÚdeviceÚdtypeÚweightÚbiasÚrunning_meanÚrunning_varÚnum_batches_trackedr!   ©r   )ÚsuperÚ__init__r   r   r   r   r   r   ÚtorchÚemptyr"   r#   Úregister_parameterÚregister_bufferÚzerosÚonesÚtensorÚlongÚitemsÚreset_parameters)Úselfr   r   r   r   r   r    r!   Úfactory_kwargsÚkÚvÚ	__class__s              €ÚW/home/mande/repo/quber/.venv/lib/python3.13/site-packages/torch/nn/modules/batchnorm.pyr)   Ú_NormBase.__init__&   s‡  ø€ ð %+Ñ;ˆÜ‰ÑÔØ(ÔØŒØ ŒØŒØ#6Ô Ø�;�;Ü#¤E§K¢K°Ñ$OÀÑ$OÓPˆDŒKÜ!¤%§+¢+¨lÑ"M¸nÑ"MÓNˆD�Ià×#Ñ# H¨dÔ3Ø×#Ñ# F¨DÔ1Ø×#×#Ø× Ñ Ø¤§¢¨LÑ K¸NÑ Kôð × Ñ ØœuŸzšz¨,ÑI¸.ÑIôñ ÙØ× Ñ Ø%Ü—’ØñäŸ*™*ðð )7×(<Ñ(<Ô(>ÔOÒ(>¡ À!ÀwÁ,“t�q’tÑ(>ÒOñ	ôò à× Ñ  °Ô6Ø× Ñ  °Ô5Ø× Ñ Ð!6¸Ô=Ø×ÑÕùó Ps   ÅGÅ'Gc                 óÆ   • U R                   (       aP  U R                  R                  5         U R                  R	                  S5        U R
                  R                  5         g g )Nr
   )r   r$   Úzero_r%   Úfill_r&   ©r4   s    r9   Úreset_running_statsÚ_NormBase.reset_running_statsV   sJ   € Ø×#×#ð ×Ñ×#Ñ#Ô%Ø×Ñ×"Ñ" 1Ô%Ø×$Ñ$×*Ñ*Õ,ð $ó    c                 óÈ   • U R                  5         U R                  (       aA  [        R                  " U R                  5        [        R
                  " U R                  5        g g ©N)r?   r   r   Úones_r"   Úzeros_r#   r>   s    r9   r3   Ú_NormBase.reset_parameters^   s:   € Ø× Ñ Ô"Ø�;�;Ü�JŠJ�t—{‘{Ô#Ü�KŠK˜Ÿ	™	Õ"ð rA   c                 ó   • [         erC   )ÚNotImplementedError©r4   Úinputs     r9   Ú_check_input_dimÚ_NormBase._check_input_dimd   s   € Ü!Ð!rA   c                 ó:   • SR                   " S0 U R                  D6$ )Nzj{num_features}, eps={eps}, momentum={momentum}, affine={affine}, track_running_stats={track_running_stats}© )ÚformatÚ__dict__r>   s    r9   Ú
extra_reprÚ_NormBase.extra_reprg   s)   € ð8ß8>¹ð?ñ PØAEÇÁñPð	
rA   c           	      ót  >• UR                  SS 5      nUb  US:  a‡  U R                  (       av  US-   n	X‘;  al  U R                  b:  U R                  R                  [        R                  " S5      :w  a  U R                  O"[        R
                  " S[        R                  S9X'   [        T
U ]!  UUUUUUU5        g )NÚversionr   r&   Úmetar   )r!   )	Úgetr   r&   r    r*   r0   r1   r(   Ú_load_from_state_dict)r4   Ú
state_dictÚprefixÚlocal_metadataÚstrictÚmissing_keysÚunexpected_keysÚ
error_msgsrT   Únum_batches_tracked_keyr8   s             €r9   rW   Ú_NormBase._load_from_state_dictm   s´   ø€ ð !×$Ñ$ Y°Ó5ˆà‰O˜w¨›{°×0H×0Hð '-Ð/DÑ&DÐ#Ø&Ó8ð ×/Ñ/Ñ;Ø×0Ñ0×7Ñ7¼5¿<º<ÈÓ;OÓOð ×,Ò,ô Ÿš a¬u¯z©zÑ:ð	 Ñ3ô 	‰Ñ%ØØØØØØØõ	
rA   )r   r#   r   r   r   r   r"   ©çñhãˆµøä>çš™™™™™¹?TTNN©r   N)Ú__name__Ú
__module__Ú__qualname__Ú__firstlineno__Ú__doc__Ú_versionÚ__constants__ÚintÚ__annotations__ÚfloatÚboolr)   r?   r3   rK   rQ   rW   Ú__static_attributes__Ú__classcell__©r8   s   @r9   r   r      sµ   ø‡ Ù6à€HÚX€MØÓØ	ƒJØ�d‰lÓØƒLØÓð Ø!$ØØ$(ØØñ. àð. ð ð. ð ˜$‘,ð	. ð
 ð. ð "ð. ð 
÷. ð . ô`-ô#ò"ò
ð 
ð 
÷ 
õ  
rA   r   c                   ól   ^ • \ rS rSr      SS\S\S\S-  S\S\SS4U 4S	 jjjrS
\S\4S jr	Sr
U =r$ )Ú
_BatchNormé�   Nr   r   r   r   r   r   c                 ó4   >• XgS.n[         T	U ]  " XX4U40 UD6  g ©Nr   )r(   r)   )
r4   r   r   r   r   r   r    r!   r5   r8   s
            €r9   r)   Ú_BatchNorm.__init__‘   s*   ø€ ð %+Ñ;ˆÜ‰ÒØ˜xÐ1Dñ	
ØHVó	
rA   rJ   c           
      óô  • U R                  U5        U R                  c  SnOU R                  nU R                  (       ak  U R                  (       aZ  U R                  bM  U R                  R                  S5        U R                  c  S[        U R                  5      -  nOU R                  n U R                  (       a  SnO#U R                  S L =(       a    U R                  S L n [        R                  " UU R                  (       a  U R                  (       a  U R                  OS U R                  (       a  U R                  (       a  U R                  OS U R                  U R                  UUU R                  5      $ )Nç        r
   ç      ð?T)rK   r   Útrainingr   r&   Úadd_rn   r$   r%   ÚFÚ
batch_normr"   r#   r   )r4   rJ   Úexponential_average_factorÚbn_trainings       r9   ÚforwardÚ_BatchNorm.forward    s'  € Ø×Ñ˜eÔ$ð
 �=‰=Ñ Ø),Ñ&à)-¯©Ð&à�=�=˜T×5×5à×'Ñ'Ñ3Ø×(Ñ(×-Ñ-¨aÔ0Ø—=‘=Ñ(Ø14´u¸T×=UÑ=UÓ7VÑ1VÑ.à15·±Ð.ð	ð �=�=Ø‰Kà×,Ñ,°Ð4×T¸4×;KÑ;KÈtÐ;SˆKð	ô
 �|Š|Øð —}—}¨×(@×(@ð ×!Ò!àà$(§M§M°T×5M×5MˆD×ÒÐSWØ�K‰KØ�I‰IØØ&Ø�H‰Hó
ð 	
rA   rN   ra   )re   rf   rg   rh   rl   rn   ro   r)   r   r‚   rp   rq   rr   s   @r9   rt   rt   �   sw   ø† ð Ø!$ØØ$(ØØñ
àð
ð ð
ð ˜$‘,ð	
ð
 ð
ð "ð
ð 
÷
ð 
ð0
˜Vð 0
¨÷ 0
ò 0
rA   rt   c                   ón   ^ • \ rS rSr% \\S'   \\S'         S S	U 4S jjjrS	U 4S jjrS	S jrSr	U =r
$ )
Ú_LazyNormBaseéÓ   r"   r#   c           
      óÐ  >• XVS.n[         T
U ]  " SUUSS40 UD6  X0l        X@l        U R                  (       a   [	        S0 UD6U l        [	        S0 UD6U l        U R                  (       ax  [        S0 UD6U l        [        S0 UD6U l	        [        R                  "  SS[        R                  0UR                  5        VV	s0 s H  u  p‰US:w  d  M  X‰_M     sn	nD6U l        g g s  sn	nf )Nr   r   Fr!   rN   r'   )r(   r)   r   r   r	   r"   r#   r   r$   r%   r*   r0   r1   r2   r&   )r4   r   r   r   r   r    r!   r5   r6   r7   r8   s             €r9   r)   Ú_LazyNormBase.__init__×   së   ø€ ð %+Ñ;ˆä‰Òð ØØØØñ		
ð ò		
ð ŒØ#6Ô Ø�;�;ä0ÑB°>ÑBˆDŒKä.Ñ@°Ñ@ˆDŒIØ×#×#ä 3Ñ E°nÑ EˆDÔä2ÑD°^ÑDˆDÔÜ',§|¢|Øñ(ä—j‘jð(ð %3×$8Ñ$8Ô$:ÔKÒ$:™D˜A¸aÀ7¹l“4�1’4Ñ$:ÒKñ	(ˆDÕ$ð $ùó Ls   Â?C"ÃC"c                 óp   >• U R                  5       (       d   U R                  S:w  a  [        TU ]  5         g g g )Nr   )Úhas_uninitialized_paramsr   r(   r3   )r4   r8   s    €r9   r3   Ú_LazyNormBase.reset_parametersÿ   s3   ø€ à×,Ñ,×.Ñ.°4×3DÑ3DÈÓ3IÜ‰GÑ$Õ&ð 4JÐ.rA   c                 ó–  • U R                  5       (       Ga3  UR                  S   U l        U R                  (       a   [	        U R
                  [        5      (       d  [        S5      e[	        U R                  [        5      (       d  [        S5      eU R
                  R                  U R                  45        U R                  R                  U R                  45        U R                  (       aL  U R                  R                  U R                  45        U R                  R                  U R                  45        U R                  5         g g )Nr
   z-self.weight must be an UninitializedParameterz+self.bias must be an UninitializedParameter)rŠ   Úshaper   r   Ú
isinstancer"   r	   ÚAssertionErrorr#   Úmaterializer   r$   r%   r3   rI   s     r9   Úinitialize_parametersÚ#_LazyNormBase.initialize_parameters  sþ   € à×(Ñ(×*Ò*Ø %§¡¨A¡ˆDÔØ�{�{Ü! $§+¡+Ô/E×FÑFÜ(ØGóð ô " $§)¡)Ô-C×DÑDÜ(Ð)VÓWÐWØ—‘×'Ñ'¨×):Ñ):Ð(<Ô=Ø—	‘	×%Ñ% t×'8Ñ'8Ð&:Ô;Ø×'×'Ø×!Ñ!×-Ñ-Ø×&Ñ&Ð(ôð × Ñ ×,Ñ,Ø×&Ñ&Ð(ôð ×!Ñ!Õ#ð% +rA   )r   r#   r&   r   r$   r%   r   r"   ra   rd   )re   rf   rg   rh   r	   rm   r)   r3   r‘   rp   rq   rr   s   @r9   r…   r…   Ó   sG   ø‡ Ø"Ó"Ø
 Ó ð ØØØ ØØð&ð 
÷&ð &÷P'÷
$ò $rA   r…   c                   ó"   • \ rS rSrSrSS jrSrg)r   i  a  Applies Batch Normalization over a 2D or 3D input.

Method described in the paper
`Batch Normalization: Accelerating Deep Network Training by Reducing
Internal Covariate Shift <https://arxiv.org/abs/1502.03167>`__ .

.. math::

    y = \frac{x - \mathrm{E}[x]}{\sqrt{\mathrm{Var}[x] + \epsilon}} * \gamma + \beta

The mean and standard-deviation are calculated per-dimension over
the mini-batches and :math:`\gamma` and :math:`\beta` are learnable parameter vectors
of size `C` (where `C` is the number of features or channels of the input). By default, the
elements of :math:`\gamma` are set to 1 and the elements of :math:`\beta` are set to 0.
At train time in the forward pass, the variance is calculated via the biased estimator,
equivalent to ``torch.var(input, correction=0)``. However, the value stored in the
moving average of the variance is calculated via the unbiased  estimator, equivalent to
``torch.var(input, correction=1)``.

Also by default, during training this layer keeps running estimates of its
computed mean and variance, which are then used for normalization during
evaluation. The running estimates are kept with a default :attr:`momentum`
of 0.1.

If :attr:`track_running_stats` is set to ``False``, this layer then does not
keep running estimates, and batch statistics are instead used during
evaluation time as well.

.. note::
    This :attr:`momentum` argument is different from one used in optimizer
    classes and the conventional notion of momentum. Mathematically, the
    update rule for running statistics here is
    :math:`\hat{x}_\text{new} = (1 - \text{momentum}) \times \hat{x} + \text{momentum} \times x_t`,
    where :math:`\hat{x}` is the estimated statistic and :math:`x_t` is the
    new observed value.

Because the Batch Normalization is done over the `C` dimension, computing statistics
on `(N, L)` slices, it's common terminology to call this Temporal Batch Normalization.

Args:
    num_features: number of features or channels :math:`C` of the input
    eps: a value added to the denominator for numerical stability.
        Default: 1e-5
    momentum: the value used for the running_mean and running_var
        computation. Can be set to ``None`` for cumulative moving average
        (i.e. simple average). Default: 0.1
    affine: a boolean value that when set to ``True``, this module has
        learnable affine parameters. Default: ``True``
    track_running_stats: a boolean value that when set to ``True``, this
        module tracks the running mean and variance, and when set to ``False``,
        this module does not track such statistics, and initializes statistics
        buffers :attr:`running_mean` and :attr:`running_var` as ``None``.
        When these buffers are ``None``, this module always uses batch statistics.
        in both training and eval modes. Default: ``True``

Shape:
    - Input: :math:`(N, C)` or :math:`(N, C, L)`, where :math:`N` is the batch size,
      :math:`C` is the number of features or channels, and :math:`L` is the sequence length
    - Output: :math:`(N, C)` or :math:`(N, C, L)` (same shape as input)

Examples::

    >>> # With Learnable Parameters
    >>> m = nn.BatchNorm1d(100)
    >>> # Without Learnable Parameters
    >>> m = nn.BatchNorm1d(100, affine=False)
    >>> input = torch.randn(20, 100)
    >>> output = m(input)
Nc                 ó�   • UR                  5       S:w  a2  UR                  5       S:w  a  [        SUR                  5        S35      eg g ©Nr   é   zexpected 2D or 3D input (got úD input)©ÚdimÚ
ValueErrorrI   s     r9   rK   ÚBatchNorm1d._check_input_dimb  ó@   € Ø�9‰9‹;˜!Ó §	¡	£¨qÓ 0ÜÐ<¸U¿Y¹Y»[¸MÈÐRÓSÐSð !1ÐrA   rN   rd   ©re   rf   rg   rh   ri   rK   rp   rN   rA   r9   r   r     s   † ñD÷LTrA   r   c                   ó&   • \ rS rSrSr\rSS jrSrg)r   ig  aþ  A :class:`torch.nn.BatchNorm1d` module with lazy initialization.

Lazy initialization based on the ``num_features`` argument of the :class:`BatchNorm1d` that is inferred
from the ``input.size(1)``.
The attributes that will be lazily initialized are `weight`, `bias`,
`running_mean` and `running_var`.

Check the :class:`torch.nn.modules.lazy.LazyModuleMixin` for further documentation
on lazy modules and their limitations.

Args:
    eps: a value added to the denominator for numerical stability.
        Default: 1e-5
    momentum: the value used for the running_mean and running_var
        computation. Can be set to ``None`` for cumulative moving average
        (i.e. simple average). Default: 0.1
    affine: a boolean value that when set to ``True``, this module has
        learnable affine parameters. Default: ``True``
    track_running_stats: a boolean value that when set to ``True``, this
        module tracks the running mean and variance, and when set to ``False``,
        this module does not track such statistics, and initializes statistics
        buffers :attr:`running_mean` and :attr:`running_var` as ``None``.
        When these buffers are ``None``, this module always uses batch statistics.
        in both training and eval modes. Default: ``True``
Nc                 ó�   • UR                  5       S:w  a2  UR                  5       S:w  a  [        SUR                  5        S35      eg g r•   r˜   rI   s     r9   rK   Ú LazyBatchNorm1d._check_input_dim„  rœ   rA   rN   rd   )	re   rf   rg   rh   ri   r   Úcls_to_becomerK   rp   rN   rA   r9   r   r   g  s   † ñð4  €M÷TrA   r   c                   ó"   • \ rS rSrSrSS jrSrg)r   i‰  aµ  Applies Batch Normalization over a 4D input.

4D is a mini-batch of 2D inputs
with additional channel dimension. Method described in the paper
`Batch Normalization: Accelerating Deep Network Training by Reducing
Internal Covariate Shift <https://arxiv.org/abs/1502.03167>`__ .

.. math::

    y = \frac{x - \mathrm{E}[x]}{ \sqrt{\mathrm{Var}[x] + \epsilon}} * \gamma + \beta

The mean and standard-deviation are calculated per-dimension over
the mini-batches and :math:`\gamma` and :math:`\beta` are learnable parameter vectors
of size `C` (where `C` is the input size). By default, the elements of :math:`\gamma` are set
to 1 and the elements of :math:`\beta` are set to 0. At train time in the forward pass, the
standard-deviation is calculated via the biased estimator, equivalent to
``torch.var(input, correction=0)``. However, the value stored in the moving average of the
standard-deviation is calculated via the unbiased  estimator, equivalent to
``torch.var(input, correction=1)``.

Also by default, during training this layer keeps running estimates of its
computed mean and variance, which are then used for normalization during
evaluation. The running estimates are kept with a default :attr:`momentum`
of 0.1.

If :attr:`track_running_stats` is set to ``False``, this layer then does not
keep running estimates, and batch statistics are instead used during
evaluation time as well.

.. note::
    This :attr:`momentum` argument is different from one used in optimizer
    classes and the conventional notion of momentum. Mathematically, the
    update rule for running statistics here is
    :math:`\hat{x}_\text{new} = (1 - \text{momentum}) \times \hat{x} + \text{momentum} \times x_t`,
    where :math:`\hat{x}` is the estimated statistic and :math:`x_t` is the
    new observed value.

Because the Batch Normalization is done over the `C` dimension, computing statistics
on `(N, H, W)` slices, it's common terminology to call this Spatial Batch Normalization.

Args:
    num_features: :math:`C` from an expected input of size
        :math:`(N, C, H, W)`
    eps: a value added to the denominator for numerical stability.
        Default: 1e-5
    momentum: the value used for the running_mean and running_var
        computation. Can be set to ``None`` for cumulative moving average
        (i.e. simple average). Default: 0.1
    affine: a boolean value that when set to ``True``, this module has
        learnable affine parameters. Default: ``True``
    track_running_stats: a boolean value that when set to ``True``, this
        module tracks the running mean and variance, and when set to ``False``,
        this module does not track such statistics, and initializes statistics
        buffers :attr:`running_mean` and :attr:`running_var` as ``None``.
        When these buffers are ``None``, this module always uses batch statistics.
        in both training and eval modes. Default: ``True``

Shape:
    - Input: :math:`(N, C, H, W)`
    - Output: :math:`(N, C, H, W)` (same shape as input)

Examples::

    >>> # With Learnable Parameters
    >>> m = nn.BatchNorm2d(100)
    >>> # Without Learnable Parameters
    >>> m = nn.BatchNorm2d(100, affine=False)
    >>> input = torch.randn(20, 100, 35, 45)
    >>> output = m(input)
Nc                 óf   • UR                  5       S:w  a  [        SUR                  5        S35      eg ©Né   zexpected 4D input (got r—   r˜   rI   s     r9   rK   ÚBatchNorm2d._check_input_dimÑ  ó0   € Ø�9‰9‹;˜!ÓÜÐ6°u·y±y³{°mÀ8ÐLÓMÐMð rA   rN   rd   r�   rN   rA   r9   r   r   ‰  ó   † ñE÷NNrA   r   c                   ó&   • \ rS rSrSr\rSS jrSrg)r   iÖ  a  A :class:`torch.nn.BatchNorm2d` module with lazy initialization.

Lazy initialization is done for the ``num_features`` argument of the :class:`BatchNorm2d` that is inferred
from the ``input.size(1)``.
The attributes that will be lazily initialized are `weight`, `bias`,
`running_mean` and `running_var`.

Check the :class:`torch.nn.modules.lazy.LazyModuleMixin` for further documentation
on lazy modules and their limitations.

Args:
    eps: a value added to the denominator for numerical stability.
        Default: 1e-5
    momentum: the value used for the running_mean and running_var
        computation. Can be set to ``None`` for cumulative moving average
        (i.e. simple average). Default: 0.1
    affine: a boolean value that when set to ``True``, this module has
        learnable affine parameters. Default: ``True``
    track_running_stats: a boolean value that when set to ``True``, this
        module tracks the running mean and variance, and when set to ``False``,
        this module does not track such statistics, and initializes statistics
        buffers :attr:`running_mean` and :attr:`running_var` as ``None``.
        When these buffers are ``None``, this module always uses batch statistics.
        in both training and eval modes. Default: ``True``
Nc                 óf   • UR                  5       S:w  a  [        SUR                  5        S35      eg r¤   r˜   rI   s     r9   rK   Ú LazyBatchNorm2d._check_input_dimó  r§   rA   rN   rd   )	re   rf   rg   rh   ri   r   r¡   rK   rp   rN   rA   r9   r   r   Ö  ó   † ñð4  €M÷NrA   r   c                   ó"   • \ rS rSrSrSS jrSrg)r   iø  aê  Applies Batch Normalization over a 5D input.

5D is a mini-batch of 3D inputs with additional channel dimension as described in the paper
`Batch Normalization: Accelerating Deep Network Training by Reducing
Internal Covariate Shift <https://arxiv.org/abs/1502.03167>`__ .

.. math::

    y = \frac{x - \mathrm{E}[x]}{ \sqrt{\mathrm{Var}[x] + \epsilon}} * \gamma + \beta

The mean and standard-deviation are calculated per-dimension over
the mini-batches and :math:`\gamma` and :math:`\beta` are learnable parameter vectors
of size `C` (where `C` is the input size). By default, the elements of :math:`\gamma` are set
to 1 and the elements of :math:`\beta` are set to 0. At train time in the forward pass, the
standard-deviation is calculated via the biased estimator, equivalent to
``torch.var(input, correction=0)``. However, the value stored in the moving average of the
standard-deviation is calculated via the unbiased  estimator, equivalent to
``torch.var(input, correction=1)``.

Also by default, during training this layer keeps running estimates of its
computed mean and variance, which are then used for normalization during
evaluation. The running estimates are kept with a default :attr:`momentum`
of 0.1.

If :attr:`track_running_stats` is set to ``False``, this layer then does not
keep running estimates, and batch statistics are instead used during
evaluation time as well.

.. note::
    This :attr:`momentum` argument is different from one used in optimizer
    classes and the conventional notion of momentum. Mathematically, the
    update rule for running statistics here is
    :math:`\hat{x}_\text{new} = (1 - \text{momentum}) \times \hat{x} + \text{momentum} \times x_t`,
    where :math:`\hat{x}` is the estimated statistic and :math:`x_t` is the
    new observed value.

Because the Batch Normalization is done over the `C` dimension, computing statistics
on `(N, D, H, W)` slices, it's common terminology to call this Volumetric Batch Normalization
or Spatio-temporal Batch Normalization.

Args:
    num_features: :math:`C` from an expected input of size
        :math:`(N, C, D, H, W)`
    eps: a value added to the denominator for numerical stability.
        Default: 1e-5
    momentum: the value used for the running_mean and running_var
        computation. Can be set to ``None`` for cumulative moving average
        (i.e. simple average). Default: 0.1
    affine: a boolean value that when set to ``True``, this module has
        learnable affine parameters. Default: ``True``
    track_running_stats: a boolean value that when set to ``True``, this
        module tracks the running mean and variance, and when set to ``False``,
        this module does not track such statistics, and initializes statistics
        buffers :attr:`running_mean` and :attr:`running_var` as ``None``.
        When these buffers are ``None``, this module always uses batch statistics.
        in both training and eval modes. Default: ``True``

Shape:
    - Input: :math:`(N, C, D, H, W)`
    - Output: :math:`(N, C, D, H, W)` (same shape as input)

Examples::

    >>> # With Learnable Parameters
    >>> m = nn.BatchNorm3d(100)
    >>> # Without Learnable Parameters
    >>> m = nn.BatchNorm3d(100, affine=False)
    >>> input = torch.randn(20, 100, 35, 45, 10)
    >>> output = m(input)
Nc                 óf   • UR                  5       S:w  a  [        SUR                  5        S35      eg ©Né   zexpected 5D input (got r—   r˜   rI   s     r9   rK   ÚBatchNorm3d._check_input_dim@  r§   rA   rN   rd   r�   rN   rA   r9   r   r   ø  r¨   rA   r   c                   ó&   • \ rS rSrSr\rSS jrSrg)r   iE  a  A :class:`torch.nn.BatchNorm3d` module with lazy initialization.

Lazy initialization is done for the ``num_features`` argument of the :class:`BatchNorm3d` that is inferred
from the ``input.size(1)``.
The attributes that will be lazily initialized are `weight`, `bias`,
`running_mean` and `running_var`.

Check the :class:`torch.nn.modules.lazy.LazyModuleMixin` for further documentation
on lazy modules and their limitations.

Args:
    eps: a value added to the denominator for numerical stability.
        Default: 1e-5
    momentum: the value used for the running_mean and running_var
        computation. Can be set to ``None`` for cumulative moving average
        (i.e. simple average). Default: 0.1
    affine: a boolean value that when set to ``True``, this module has
        learnable affine parameters. Default: ``True``
    track_running_stats: a boolean value that when set to ``True``, this
        module tracks the running mean and variance, and when set to ``False``,
        this module does not track such statistics, and initializes statistics
        buffers :attr:`running_mean` and :attr:`running_var` as ``None``.
        When these buffers are ``None``, this module always uses batch statistics.
        in both training and eval modes. Default: ``True``
Nc                 óf   • UR                  5       S:w  a  [        SUR                  5        S35      eg r¯   r˜   rI   s     r9   rK   Ú LazyBatchNorm3d._check_input_dimb  r§   rA   rN   rd   )	re   rf   rg   rh   ri   r   r¡   rK   rp   rN   rA   r9   r   r   E  r¬   rA   r   c                   ó¤   ^ • \ rS rSrSr       SS\S\S\S-  S\S\S	\S-  S
S4U 4S jjjr	SS jr
SS jrS\S
\4S jr\SS j5       rSrU =r$ )r   ig  a¼  Applies Batch Normalization over a N-Dimensional input.

The N-D input is a mini-batch of [N-2]D inputs with additional channel dimension) as described in the paper
`Batch Normalization: Accelerating Deep Network Training by Reducing
Internal Covariate Shift <https://arxiv.org/abs/1502.03167>`__ .

.. math::

    y = \frac{x - \mathrm{E}[x]}{ \sqrt{\mathrm{Var}[x] + \epsilon}} * \gamma + \beta

The mean and standard-deviation are calculated per-dimension over all
mini-batches of the same process groups. :math:`\gamma` and :math:`\beta`
are learnable parameter vectors of size `C` (where `C` is the input size).
By default, the elements of :math:`\gamma` are sampled from
:math:`\mathcal{U}(0, 1)` and the elements of :math:`\beta` are set to 0.
The standard-deviation is calculated via the biased estimator, equivalent to
`torch.var(input, correction=0)`.

Also by default, during training this layer keeps running estimates of its
computed mean and variance, which are then used for normalization during
evaluation. The running estimates are kept with a default :attr:`momentum`
of 0.1.

If :attr:`track_running_stats` is set to ``False``, this layer then does not
keep running estimates, and batch statistics are instead used during
evaluation time as well.

.. note::
    This :attr:`momentum` argument is different from one used in optimizer
    classes and the conventional notion of momentum. Mathematically, the
    update rule for running statistics here is
    :math:`\hat{x}_\text{new} = (1 - \text{momentum}) \times \hat{x} + \text{momentum} \times x_t`,
    where :math:`\hat{x}` is the estimated statistic and :math:`x_t` is the
    new observed value.

Because the Batch Normalization is done for each channel in the ``C`` dimension, computing
statistics on ``(N, +)`` slices, it's common terminology to call this Volumetric Batch
Normalization or Spatio-temporal Batch Normalization.

Currently :class:`SyncBatchNorm` only supports
:class:`~torch.nn.DistributedDataParallel` (DDP) with single GPU per process. Use
:meth:`torch.nn.SyncBatchNorm.convert_sync_batchnorm()` to convert
:attr:`BatchNorm*D` layer to :class:`SyncBatchNorm` before wrapping
Network with DDP.

Args:
    num_features: :math:`C` from an expected input of size
        :math:`(N, C, +)`
    eps: a value added to the denominator for numerical stability.
        Default: ``1e-5``
    momentum: the value used for the running_mean and running_var
        computation. Can be set to ``None`` for cumulative moving average
        (i.e. simple average). Default: 0.1
    affine: a boolean value that when set to ``True``, this module has
        learnable affine parameters. Default: ``True``
    track_running_stats: a boolean value that when set to ``True``, this
        module tracks the running mean and variance, and when set to ``False``,
        this module does not track such statistics, and initializes statistics
        buffers :attr:`running_mean` and :attr:`running_var` as ``None``.
        When these buffers are ``None``, this module always uses batch statistics.
        in both training and eval modes. Default: ``True``
    process_group: synchronization of stats happen within each process group
        individually. Default behavior is synchronization across the whole
        world

Shape:
    - Input: :math:`(N, C, +)`
    - Output: :math:`(N, C, +)` (same shape as input)

.. note::
    Synchronization of batchnorm statistics occurs only while training, i.e.
    synchronization is disabled when ``model.eval()`` is set or if
    ``self.training`` is otherwise ``False``.

Examples::

    >>> # xdoctest: +SKIP
    >>> # With Learnable Parameters
    >>> m = nn.SyncBatchNorm(100)
    >>> # creating process group (optional)
    >>> # ranks is a list of int identifying rank ids.
    >>> ranks = list(range(8))
    >>> r1, r2 = ranks[:4], ranks[4:]
    >>> # Note: every rank calls into new_group for every
    >>> # process group created, even if that rank is not
    >>> # part of the group.
    >>> process_groups = [torch.distributed.new_group(pids) for pids in [r1, r2]]
    >>> process_group = process_groups[0 if dist.get_rank() <= 3 else 1]
    >>> # Without Learnable Parameters
    >>> m = nn.BatchNorm3d(100, affine=False, process_group=process_group)
    >>> input = torch.randn(20, 100, 35, 45, 10)
    >>> output = m(input)

    >>> # network is nn.BatchNorm layer
    >>> sync_bn_network = nn.SyncBatchNorm.convert_sync_batchnorm(network, process_group)
    >>> # only single gpu per process is currently supported
    >>> ddp_sync_bn_network = torch.nn.parallel.DistributedDataParallel(
    >>>                         sync_bn_network,
    >>>                         device_ids=[args.local_rank],
    >>>                         output_device=args.local_rank)
Nr   r   r   r   r   Úprocess_groupr   c	                 ó@   >• XxS.n	[         T
U ]  " XX4U40 U	D6  X`l        g rw   )r(   r)   r¶   )r4   r   r   r   r   r   r¶   r    r!   r5   r8   s             €r9   r)   ÚSyncBatchNorm.__init__Î  s2   ø€ ð %+Ñ;ˆÜ‰ÒØ˜xÐ1Dñ	
ØHVò	
ð +ÕrA   c                 óf   • UR                  5       S:  a  [        SUR                  5        S35      eg )Nr   z expected at least 2D input (got r—   r˜   rI   s     r9   rK   ÚSyncBatchNorm._check_input_dimß  s/   € Ø�9‰9‹;˜‹?ÜÐ?ÀÇ	Á	Ã¸}ÈHÐUÓVÐVð rA   c                 óD   • UR                  S5      S:X  a  [        S5      eg )Nr
   r   z9SyncBatchNorm number of input channels should be non-zero)Úsizerš   rI   s     r9   Ú_check_non_zero_input_channelsÚ,SyncBatchNorm._check_non_zero_input_channelsã  s'   € Ø�:‰:�a‹=˜AÓÜØKóð ð rA   rJ   c                 óF  • U R                  U5        U R                  U5        U R                  c  SnOU R                  nU R                  (       a{  U R                  (       aj  U R
                  c  [        S5      eU R
                  R                  S5        U R                  c  SU R
                  R                  5       -  nOU R                  n U R                  (       a  SnO#U R                  SL =(       a    U R                  SL n U R                  (       a  U R                  (       a  U R                  OSnU R                  (       a  U R                  (       a  U R                  OSnU=(       aV    U R                  =(       aC    [        R                  R                  5       =(       a    [        R                  R                  5       nU(       aÉ  UR                  R                   SSS	[        R"                  R%                  5       4;  a*  ['        S
[        R"                  R%                  5        35      e[        R                  R(                  R*                  nU R,                  (       a  U R,                  n[        R                  R/                  U5      nUS:„  nU(       d;  [0        R2                  " UUUU R4                  U R6                  UUU R8                  5      $ U(       d  [        S5      e[:        R<                  " UU R4                  U R6                  UUU R8                  UWW5	      $ )z
Runs the forward pass.
Nrz   z$num_batches_tracked must not be Noner
   r{   TÚcudaÚhpuÚxpuz;SyncBatchNorm expected input tensor to be on GPU or XPU or zbn_training must be True)rK   r½   r   r|   r   r&   r�   r}   Úitemr$   r%   r*   ÚdistributedÚis_availableÚis_initializedr    ÚtypeÚ_CÚ_get_privateuse1_backend_namerš   ÚgroupÚWORLDr¶   Úget_world_sizer~   r   r"   r#   r   Úsync_batch_normÚapply)	r4   rJ   r€   r�   r$   r%   Ú	need_syncr¶   Ú
world_sizes	            r9   r‚   ÚSyncBatchNorm.forwardé  s–  € ð 	×Ñ˜eÔ$Ø×+Ñ+¨EÔ2ð
 �=‰=Ñ Ø),Ñ&à)-¯©Ð&à�=�=˜T×5×5Ø×'Ñ'Ñ/Ü$Ð%KÓLÐLØ×$Ñ$×)Ñ)¨!Ô,Ø�}‰}Ñ$Ø-0°4×3KÑ3K×3PÑ3PÓ3RÑ-RÑ*à-1¯]©]Ð*ð	ð �=�=Ø‰Kà×,Ñ,°Ð4×T¸4×;KÑ;KÈtÐ;SˆKð	ð &*§]§]°d×6N×6NˆD×ÒÐTXð 	ð %)§M§M°T×5M×5MˆD×ÒÐSWð 	ð ÷ 3Ø—‘÷3ä×!Ñ!×.Ñ.Ó0÷3ô ×!Ñ!×0Ñ0Ó2ð	 	ö à�|‰|× Ñ ØØØÜ—‘×6Ñ6Ó8ð	)ó ô !ØQÜ—x‘x×=Ñ=Ó?Ð@ðBóð ô
 "×-Ñ-×3Ñ3×9Ñ9ˆMØ×!×!Ø $× 2Ñ 2�Ü×*Ñ*×9Ñ9¸-ÓHˆJØ" Q™ˆIö Ü—<’<ØØØØ—‘Ø—	‘	ØØ*Ø—‘ó	ð 	ö Ü$Ð%?Ó@Ð@Ü"×(Ò(ØØ—‘Ø—	‘	ØØØ—‘Ø*ØØó
ð 
rA   c                 ó6  • Un[        U[        R                  R                  R                  R
                  5      (       Ga  [        R                  R                  UR                  UR                  UR                  UR                  UR                  U5      nUR                  (       a@  [        R                  " 5          UR                  Ul        UR                  Ul        SSS5        UR                  Ul        UR                   Ul        UR"                  Ul        UR$                  Ul        ['        US5      (       a  UR(                  Ul        UR+                  5        H%  u  pEUR-                  X@R/                  XR5      5        M'     AU$ ! , (       d  f       N°= f)a�  Converts all :attr:`BatchNorm*D` layers in the model to :class:`torch.nn.SyncBatchNorm` layers.

Args:
    module (nn.Module): module containing one or more :attr:`BatchNorm*D` layers
    process_group (optional): process group to scope synchronization,
        default is the whole world

Returns:
    The original :attr:`module` with the converted :class:`torch.nn.SyncBatchNorm`
    layers. If the original :attr:`module` is a :attr:`BatchNorm*D` layer,
    a new :class:`torch.nn.SyncBatchNorm` layer object will be returned
    instead.

Example::

    >>> # Network with nn.BatchNorm layer
    >>> # xdoctest: +REQUIRES(env:TORCH_DOCTEST_CUDA)
    >>> module = torch.nn.Sequential(
    >>>            torch.nn.Linear(20, 100),
    >>>            torch.nn.BatchNorm1d(100),
    >>>          ).cuda()
    >>> # creating process group (optional)
    >>> # ranks is a list of int identifying rank ids.
    >>> ranks = list(range(8))
    >>> r1, r2 = ranks[:4], ranks[4:]
    >>> # Note: every rank calls into new_group for every
    >>> # process group created, even if that rank is not
    >>> # part of the group.
    >>> # xdoctest: +SKIP("distributed")
    >>> process_groups = [torch.distributed.new_group(pids) for pids in [r1, r2]]
    >>> process_group = process_groups[0 if dist.get_rank() <= 3 else 1]
    >>> sync_bn_module = torch.nn.SyncBatchNorm.convert_sync_batchnorm(module, process_group)

NÚqconfig)rŽ   r*   ÚnnÚmodulesÚ	batchnormrt   r   r   r   r   r   r   Úno_gradr"   r#   r$   r%   r&   r|   ÚhasattrrÓ   Únamed_childrenÚ
add_moduleÚconvert_sync_batchnorm)ÚclsÚmoduler¶   Úmodule_outputÚnameÚchilds         r9   rÛ   Ú$SyncBatchNorm.convert_sync_batchnormL  s/  € ðH ˆÜ�fœeŸh™h×.Ñ.×8Ñ8×CÑC×DÒDÜ!ŸH™H×2Ñ2Ø×#Ñ#Ø—
‘
Ø—‘Ø—‘Ø×*Ñ*ØóˆMð �}�}Ü—]’]•_Ø+1¯=©=�MÔ(Ø)/¯©�MÔ&÷ %ð *0×)<Ñ)<ˆMÔ&Ø(.×(:Ñ(:ˆMÔ%Ø06×0JÑ0JˆMÔ-Ø%+§_¡_ˆMÔ"Ü�v˜y×)Ñ)Ø(.¯©�Ô%Ø!×0Ñ0Ö2‰KˆDØ×$Ñ$Ø×0Ñ0°ÓFöñ 3ð ØÐ÷ %•_ús   Â=#F
Æ

F)r¶   )rb   rc   TTNNNrd   rC   )re   rf   rg   rh   ri   rl   rn   ro   r   r)   rK   r½   r   r‚   ÚclassmethodrÛ   rp   rq   rr   s   @r9   r   r   g  s­   ø† ñdðR Ø!$ØØ$(Ø$(ØØñ+àð+ð ð+ð ˜$‘,ð	+ð
 ð+ð "ð+ð ˜T‘zð+ð 
÷+ð +ô"Wôða˜Vð a¨ô aðF ó<ó ö<rA   r   )Útypingr   r*   r   Útorch.nnr   r~   r   Útorch.nn.parameterr   r   r	   Ú
_functionsr   rÍ   Úlazyr   rÝ   r   Ú__all__r   rt   r…   r   r   r   r   r   r   rN   rA   r9   Ú<module>ré      sÊ   ðå ã Ý ß *ß UÑ Uå 8Ý !Ý ò€ôt
�ô t
ôn@
�ô @
ôFE$�O Yô E$ôPIT�*ô ITôXT�m Zô TôDJN�*ô JNôZN�m Zô NôDJN�*ô JNôZN�m Zô NôDb�Jõ brA   