+
    &jh                        R t ^ RIt^ RIt^ RIHt ^ RIHtHt ^ RIH	t	 ^ RI
t
^ RI
Ht . R6Ot]! R4      t]	! R4      t]R7,          t]R8,          tR9R R	 lltR9R
 R lltR9R R lltR R ltR R ltR9R R lltR:R R lltR:R R lltR;R R lltR R ltR R ltR R ltR  R! ltR<R" R# lltR$ R% ltR=R& R' llt R=R( R) llt!R* R+ lt"R>R, R- llt#R>R. R/ llt$R?R0 R1 llt%R@R2 R3 llt&R4 R5 lt']'! ]4      t(]'! ]4      t)]'! ]4      t*]'! ]4      t+]'! ]4      t,]'! ] 4      t-]'! ]!4      t.]'! ]#4      t/]'! ]$4      t0]'! ]%4      t1]'! ]&4      t2R# )AzHThis file contains utilities for initializing neural network parameters.N)Callable)LiteralTypeVar)	ParamSpecTensor_R_Pc          
      v    V ^8  d   QhR\         R\        R\        R\        P                  R,          R\         /#    tensorab	generatorNreturnr   floattorch	Generator)formats   "e/Users/jameslopez/projects/CWCArchive/cwc-podcast/.venv/lib/python3.14/site-packages/torch/nn/init.py__annotate__r   E   s=     : :::!&:38??T3I::    c                     \         P                  ! 4       ;_uu_ 4        V P                  WVR 7      uuRRR4       #   + '       g   i     R# ; ir   N)r   no_graduniform_r   r   r   r   s   &&&&r   _no_grad_uniform_r    E   s+     
qy9 
	   <A	c          
      v    V ^8  d   QhR\         R\        R\        R\        P                  R,          R\         /# r   r   meanstdr   Nr   r   )r   s   "r   r   r   L   sC     > >>
> 
> %	>
 >r   c                     \         P                  ! 4       ;_uu_ 4        V P                  WVR 7      uuRRR4       #   + '       g   i     R# ; ir   )r   r   normal_r   r$   r%   r   s   &&&&r   _no_grad_normal_r)   L   s+     
~~d9~= 
r!   c                    V ^8  d   QhR\         R\        R\        R\        R\        R\        P                  R,          R\         /# 	r   r   r$   r%   r   r   r   Nr   r   )r   s   "r   r   r   V   s`     J JJ
J 
J 	J
 J %J Jr   c                 <   V P                   '       d   V # R  R lpW^V,          ,
          8  g   W^V,          ,           8  d   \        P                  ! R^R7       \        P                  ! 4       ;_uu_ 4        V! WA,
          V,          4      V! W1,
          V,          4      ,
          pVR8  d   V P                  VRR7      P                  4       pV P                  VRR7      P                  4       p	V P                  WVR7      p
 W8  W8  ,          pVP                  4       '       g   M?\        P                  ! V\        P                  ! V
4      P                  WVR7      V
4      p
Kd  W
Jd   V P                  V
4       EM\        V\        W4      4      pR	W,
          V,          ^,          ,          p\        P                  ! V 4      p\        P                  ! V 4      pV P                  W4VR7       VP                  V 4       VP                  V4      P!                  V4      P#                  ^4      P%                  R	4      P                  V4       VP                  VR7      P'                  4       P)                  V4      pVP                  4       '       g   MT p
 VP                  W4VR7       \        P                  ! VW4      p
VP                  V4      P!                  V4      P#                  ^4      P%                  R	4      P                  V4       \        P                  ! VVP                  VR7      P'                  4       P)                  V4      V4      pVP                  4       '       d   K   T P                  T
4       V uuRRR4       #   + '       g   i     R# ; i)
c                0    V ^8  d   QhR\         R\         /# )r   xr   )r   )r   s   "r   r   ,_no_grad_trunc_normal_.<locals>.__annotate__b   s     : :E :e :r   c                     R \         P                  ! V \         P                  ! R4      ,          4      ,           R,          # )      ?       @)matherfsqrt)r.   s   &r   norm_cdf(_no_grad_trunc_normal_.<locals>.norm_cdfb   s(    dhhq499S>122c99r   zjmean is more than 2 std from [a, b] in nn.init.trunc_normal_. The distribution of values may be incorrect.
stacklevelg333333?cpu)devicer   Ng      )is_metawarningswarnr   r   
new_tensoritemr'   anywhere
empty_likecopy_maxminr   sub_div_pow_mul_log_gt)r   r$   r%   r   r   r   r6   plohiresultmaskmodelog_peak
candidates
accept_bufpendings   &&&&&&           r   _no_grad_trunc_normal_rW   V   s    ~~~: 	1s7{1s7{ 2;	
 
ah#%&18s2B)CCs7 ""1U"388:B""1U"388:B^^D^CF4xxzz$$V,44T)4T
 #V$q#d,'Ds2q88H))&1J))&1J OOAIO6V$OOD!&&s+00388>CCHM ))I)>CCEHHTG;;==''	'B"[[*EFOOD)..s388;@@FKKHU#kk"++i+@EEGJJ:VG
 #;;==V$i 
s   8K/N
,N

N	c                <    V ^8  d   QhR\         R\        R\         /# r   r   valr   r   r   )r   s   "r   r   r      s!     ! !6 ! !& !r   c                     \         P                  ! 4       ;_uu_ 4        V P                  V4      uuR R R 4       #   + '       g   i     R # ; iN)r   r   fill_r   rZ   s   &&r   _no_grad_fill_r`      s%    	||C  
s	   :A	c                0    V ^8  d   QhR\         R\         /# r   r   r   r   )r   s   "r   r   r      s      6 f r   c                     \         P                  ! 4       ;_uu_ 4        V P                  4       uuR R R 4       #   + '       g   i     R # ; ir]   )r   r   zero_r   s   &r   _no_grad_zero_rf      s"    	||~ 
s	   9A
	c                `    V ^8  d   QhR\         R\        \        ,          R,          R\        /# )r   nonlinearityparamNr   )_NonlinearityTypeintr   )r   s   "r   r   r      s7     GE GE#GE,/%K$,>GE
GEr   c                   . ROpW9   g   V R8X  d   ^# V R8X  d   R# V R8X  d   \         P                  ! R4      # V R8X  d   Vf   RpMT\        V\        4      '       g   \        V\        4      '       g   \        V\
        4      '       d   TpM\        RV R24      h\         P                  ! R^V^,          ,           ,          4      # V R	8X  d   R# \        R
V  24      h)a  Return the recommended gain value for the given nonlinearity function.

The values are as follows:

================= ====================================================
nonlinearity      gain
================= ====================================================
Linear / Identity :math:`1`
Conv{1,2,3}D      :math:`1`
Sigmoid           :math:`1`
Tanh              :math:`\frac{5}{3}`
ReLU              :math:`\sqrt{2}`
Leaky Relu        :math:`\sqrt{\frac{2}{1 + \text{negative\_slope}^2}}`
SELU              :math:`\frac{3}{4}`
================= ====================================================

.. warning::
    In order to implement `Self-Normalizing Neural Networks`_ ,
    you should use ``nonlinearity='linear'`` instead of ``nonlinearity='selu'``.
    This gives the initial weights a variance of ``1 / N``,
    which is necessary to induce a stable fixed point in the forward pass.
    In contrast, the default gain for ``SELU`` sacrifices the normalization
    effect for more stable gradient flow in rectangular layers.

Args:
    nonlinearity: the non-linear function (`nn.functional` name)
    param: optional parameter for the non-linear function

Examples:
    >>> gain = nn.init.calculate_gain(
    ...     "leaky_relu", 0.2
    ... )  # leaky_relu with negative_slope=0.2

.. _Self-Normalizing Neural Networks: https://papers.nips.cc/paper/2017/hash/5d44ee6f2c3f71b73125876103c8f6c4-Abstract.html
sigmoidtanhrelur2   
leaky_relu{Gz?znegative_slope z not a valid numberseluzUnsupported nonlinearity )linearconv1dconv2dconv3dconv_transpose1dconv_transpose2dconv_transpose3dg?g      ?)r3   r5   
isinstanceboolrk   r   
ValueError)rh   ri   
linear_fnsnegative_slopes   &&  r   calculate_gainr      s    LJ !\Y%>				yy~		%=!N5$''5#&&%'' #Nug5HIJJyyNA$5 5677			
 4\NCDDr   c          
      v    V ^8  d   QhR\         R\        R\        R\        P                  R,          R\         /# r   r   )r   s   "r   r   r      sC     6 666 6 %	6
 6r   c           	         \         P                  P                  V 4      '       d)   \         P                  P                  \        V 3WW#R7      # \        WW#4      # )a  Fill the input Tensor with values drawn from the uniform distribution.

:math:`\mathcal{U}(a, b)`.

Args:
    tensor: an n-dimensional `torch.Tensor`
    a: the lower bound of the uniform distribution
    b: the upper bound of the uniform distribution
    generator: the torch Generator to sample from (default: None)

Examples:
    >>> w = torch.empty(3, 5)
    >>> nn.init.uniform_(w)
r   )r   	overrideshas_torch_function_variadichandle_torch_functionr   r    r   s   &&&&r   r   r      sO    ( 226::44viq 5 
 	
 V55r   c          
      v    V ^8  d   QhR\         R\        R\        R\        P                  R,          R\         /# r#   r   )r   s   "r   r   r     sC     : ::
: 
: %	:
 :r   c           	         \         P                  P                  V 4      '       d)   \         P                  P                  \        V 3WW#R7      # \        WW#4      # )a  Fill the input Tensor with values drawn from the normal distribution.

:math:`\mathcal{N}(\text{mean}, \text{std}^2)`.

Args:
    tensor: an n-dimensional `torch.Tensor`
    mean: the mean of the normal distribution
    std: the standard deviation of the normal distribution
    generator: the torch Generator to sample from (default: None)

Examples:
    >>> w = torch.empty(3, 5)
    >>> nn.init.normal_(w)
r(   )r   r   r   r   r'   r)   r(   s   &&&&r   r'   r'     sO    ( 226::44fYvc 5 
 	
 F#99r   c                    V ^8  d   QhR\         R\        R\        R\        R\        R\        P                  R,          R\         /# r+   r   )r   s   "r   r   r   -  s`     !P !P!P
!P 
!P 	!P
 !P %!P !Pr   c           	         \        WW#WER7      # )a  Fill the input Tensor with values drawn from a truncated normal distribution.

The values are effectively drawn from the
normal distribution :math:`\mathcal{N}(\text{mean}, \text{std}^2)`
with values outside :math:`[a, b]` redrawn until they are within
the bounds. The method used for generating the random values works
best when :math:`a \leq \text{mean} \leq b`.

For reduced-precision types (``torch.float16`` and ``torch.bfloat16``),
sampling quality depends on the underlying ``normal_()`` and ``uniform_()``
implementations which operate at higher internal precision to avoid
quantization artifacts.

Args:
    tensor: an n-dimensional `torch.Tensor`
    mean: the mean of the normal distribution
    std: the standard deviation of the normal distribution
    a: the minimum cutoff value
    b: the maximum cutoff value
    generator: the torch Generator to sample from (default: None)

Examples:
    >>> w = torch.empty(3, 5)
    >>> nn.init.trunc_normal_(w)
r   )rW   )r   r$   r%   r   r   r   s   &&&&&&r   trunc_normal_r   -  s    B "&OOr   c                <    V ^8  d   QhR\         R\        R\         /# rY   r[   )r   s   "r   r   r   Q  s!     ' 'f '5 'V 'r   c                    \         P                  P                  V 4      '       d(   \         P                  P                  \        V 3WR7      # \        W4      # )zFill the input Tensor with the value :math:`\text{val}`.

Args:
    tensor: an n-dimensional `torch.Tensor`
    val: the value to fill the tensor with

Examples:
    >>> w = torch.empty(3, 5)
    >>> nn.init.constant_(w, 0.3)
r_   )r   r   r   r   	constant_r`   r_   s   &&r   r   r   Q  sK     226::44y 5 
 	
 &&&r   c                0    V ^8  d   QhR\         R\         /# rb   r   )r   s   "r   r   r   c  s     
' 
'& 
'V 
'r   c                    \        V R4      # )zFill the input Tensor with the scalar value `1`.

Args:
    tensor: an n-dimensional `torch.Tensor`

Examples:
    >>> w = torch.empty(3, 5)
    >>> nn.init.ones_(w)
r1   )r`   re   s   &r   ones_r   c  s     &#&&r   c                0    V ^8  d   QhR\         R\         /# rb   r   )r   s   "r   r   r   p  s     
" 
"6 
"f 
"r   c                    \        V 4      # )zFill the input Tensor with the scalar value `0`.

Args:
    tensor: an n-dimensional `torch.Tensor`

Examples:
    >>> w = torch.empty(3, 5)
    >>> nn.init.zeros_(w)
)rf   re   s   &r   zeros_r   p  s     &!!r   c                0    V ^8  d   QhR\         R\         /# rb   r   )r   s   "r   r   r   }  s       F r   c           	        V P                  4       ^8w  d   \        R4      h\        P                  ! 4       ;_uu_ 4        \        P                  ! V P
                  RV RV P                  /  RRR4       V #   + '       g   i     T # ; i)a  Fill the 2-dimensional input `Tensor` with the identity matrix.

Preserves the identity of the inputs in `Linear` layers, where as
many inputs are preserved as possible.

Args:
    tensor: a 2-dimensional `torch.Tensor`

Examples:
    >>> w = torch.empty(3, 5)
    >>> nn.init.eye_(w)
,Only tensors with 2 dimensions are supportedoutrequires_gradN)
ndimensionr|   r   r   eyeshaper   re   s   &r   eye_r   }  sc     aGHH			6<<PVP6;O;OP 
M 
Ms   -A55B	c                <    V ^8  d   QhR\         R\        R\         /# )r   r   groupsr   )r   rk   )r   s   "r   r   r     s!     5 56 53 5v 5r   c                   V P                  4       pVR9  d   \        R4      hV P                  4       pV^ ,          V,          ^ 8w  d   \        R4      hV P                  '       d   V # V^ ,          V,          p\	        WC^,          4      p\
        P                  ! 4       ;_uu_ 4        V P                  4        \        V4       F  p\        V4       F  pV^8X  d-   ^WV,          V,           WpP                  ^4      ^,          3&   K6  V^8X  dE   ^V Wd,          V,           VV P                  ^4      ^,          V P                  ^4      ^,          3&   K  ^V Wd,          V,           VV P                  ^4      ^,          V P                  ^4      ^,          V P                  ^4      ^,          3&   K  	  K  	  RRR4       V #   + '       g   i     T # ; i)a  Fill the {3, 4, 5}-dimensional input `Tensor` with the Dirac delta function.

Preserves the identity of the inputs in `Convolutional`
layers, where as many input channels are preserved as possible. In case
of groups>1, each group of channels preserves identity

Args:
    tensor: a {3, 4, 5}-dimensional `torch.Tensor`
    groups (int, optional): number of groups in the conv layer (default: 1)
Examples:
    >>> w = torch.empty(3, 16, 5, 5)
    >>> nn.init.dirac_(w)
    >>> w = torch.empty(3, 24, 5, 5)
    >>> nn.init.dirac_(w, 3)
z5Only tensors with 3, 4, or 5 dimensions are supportedz!dim 0 must be divisible by groupsN)         )	r   r|   sizer<   rF   r   r   rd   range)r   r   
dimensionssizesout_chans_per_grpmin_dimgds   &&      r   dirac_r     s     ""$J"PQQKKMEQx&A<==~~~aF*#1X.G	vA7^?PQF0014aQ19LLM1_  -1A!+A!+-  -1A!+A!+A!+	- $  
, M- 
, Ms   &DF<<G	c                R    V ^8  d   QhR\         R\        \        \        3,          /# rb   )r   tuplerk   )r   s   "r   r   r     s"      & U38_ r   c                 "   V P                  4       pV^8  d   \        R4      hV P                  ^4      pV P                  ^ 4      p^pV P                  4       ^8  d#   V P                  R,           F  pWE,          pK  	  W$,          pW4,          pWg3# )r   zNFan in and fan out can not be computed for tensor with fewer than 2 dimensions:r   NN)dimr|   r   r   )r   r   num_input_fmapsnum_output_fmapsreceptive_field_sizesfan_infan_outs   &       r   _calculate_fan_in_and_fan_outr     s    JA~\
 	
 kk!nO{{1~zz|a b!!A %  "3F5G?r   c                j    V ^8  d   QhR\         R\        R\        P                  R,          R\         /# r   r   gainr   Nr   r   )r   s   "r   r   r     s9     7 77
7 %7 	7r   c                    \        V 4      w  r4V\        P                  ! R\        W4,           4      ,          4      ,          p\        P                  ! R4      V,          p\	        W) Wb4      # )a  Fill the input `Tensor` with values using a Xavier uniform distribution.

The method is described in `Understanding the difficulty of training
deep feedforward neural networks` - Glorot, X. & Bengio, Y. (2010).
The resulting tensor will have values sampled from
:math:`\mathcal{U}(-a, a)` where

.. math::
    a = \text{gain} \times \sqrt{\frac{6}{\text{fan\_in} + \text{fan\_out}}}

Also known as Glorot initialization.

Args:
    tensor: an n-dimensional `torch.Tensor`
    gain: an optional scaling factor
    generator: the torch Generator to sample from (default: None)

Examples:
    >>> w = torch.empty(3, 5)
    >>> nn.init.xavier_uniform_(w, gain=nn.init.calculate_gain("relu"))
r2         @)r   r3   r5   r   r    )r   r   r   r   r   r%   r   s   &&&    r   xavier_uniform_r     sQ    4 4F;OF
3v'7!889
9C		#AVR66r   c                j    V ^8  d   QhR\         R\        R\        P                  R,          R\         /# r   r   )r   s   "r   r   r      s9     9 99
9 %9 	9r   c                    \        V 4      w  r4V\        P                  ! R\        W4,           4      ,          4      ,          p\	        V RWR4      # )a  Fill the input `Tensor` with values using a Xavier normal distribution.

The method is described in `Understanding the difficulty of training deep feedforward
neural networks` - Glorot, X. & Bengio, Y. (2010). The resulting tensor
will have values sampled from :math:`\mathcal{N}(0, \text{std}^2)` where

.. math::
    \text{std} = \text{gain} \times \sqrt{\frac{2}{\text{fan\_in} + \text{fan\_out}}}

Also known as Glorot initialization.

Args:
    tensor: an n-dimensional `torch.Tensor`
    gain: an optional scaling factor
    generator: the torch Generator to sample from (default: None)

Examples:
    >>> w = torch.empty(3, 5)
    >>> nn.init.xavier_normal_(w)
r2           )r   r3   r5   r   r)   )r   r   r   r   r   r%   s   &&&   r   xavier_normal_r      s?    2 4F;OF
3v'7!889
9CFC88r   c                <    V ^8  d   QhR\         R\        R\        /# )r   r   rR   r   )r   _FanModerk   )r   s   "r   r   r     s!     3 36 3 3c 3r   c                     VP                  4       pR R.pW9  d   \        RV RV 24      h\        V 4      w  r4VR 8X  d   V# T# )r   r   zMode z" not supported, please use one of )lowerr|   r   )r   rR   valid_modesr   r   s   &&   r   _calculate_correct_fanr     sT    ::<DY'K5&HVWW3F;OFX%6272r   c                    V ^8  d   QhR\         R\        R\        R\        R\        P
                  R,          R\         /# r   r   r   rR   rh   r   Nr   r   r   r   rj   r   r   )r   s   "r   r   r   *  sU     >C >C>C>C >C $	>C
 %>C >Cr   c           
     8   \         P                  P                  V 4      '       d,   \         P                  P                  \        V 3V VVVVR7      # ^ V P
                  9   d   \        P                  ! R^R7       V # \        W4      p\        W14      pV\        P                  ! V4      ,          p\        P                  ! R4      V,          p\         P                  ! 4       ;_uu_ 4        V P                  V) WR7      uuRRR4       #   + '       g   i     R# ; i)aD  Fill the input `Tensor` with values using a Kaiming uniform distribution.

The method is described in `Delving deep into rectifiers: Surpassing
human-level performance on ImageNet classification` - He, K. et al. (2015).
The resulting tensor will have values sampled from
:math:`\mathcal{U}(-\text{bound}, \text{bound})` where

.. math::
    \text{bound} = \text{gain} \times \sqrt{\frac{3}{\text{fan\_mode}}}

Also known as He initialization.

Args:
    tensor: an n-dimensional `torch.Tensor`
    a: the negative slope of the rectifier used after this layer (only
        used with ``'leaky_relu'``)
    mode: either ``'fan_in'`` (default) or ``'fan_out'``. Choosing ``'fan_in'``
        preserves the magnitude of the variance of the weights in the
        forward pass. Choosing ``'fan_out'`` preserves the magnitudes in the
        backwards pass.
    nonlinearity: the non-linear function (`nn.functional` name),
        recommended to use only with ``'relu'`` or ``'leaky_relu'`` (default).
    generator: the torch Generator to sample from (default: None)

Examples:
    >>> w = torch.empty(3, 5)
    >>> nn.init.kaiming_uniform_(w, mode="fan_in", nonlinearity="relu")

Note:
    Be aware that ``fan_in`` and ``fan_out`` are calculated assuming
    that the weight matrix is used in a transposed manner,
    (i.e., ``x @ w.T`` in ``Linear`` layers, where ``w.shape = [fan_out, fan_in]``).
    This is important for correct initialization.
    If you plan to use ``x @ w``, where ``w.shape = [fan_in, fan_out]``,
    pass in a transposed weight matrix, i.e. ``nn.init.kaiming_uniform_(w.T, ...)``.
)r   r   rR   rh   r   ,Initializing zero-element tensors is a no-opr8   r   r   N)r   r   r   r   kaiming_uniform_r   r=   r>   r   r   r3   r5   r   r   )	r   r   rR   rh   r   fanr   r%   bounds	   &&&&&    r   r   r   *  s    V 226::44I% 5 
 	
 	FLLDQRS
 
.C,*D
3
CIIcNS E	vuB 
s   )DD	c                    V ^8  d   QhR\         R\        R\        R\        R\        P
                  R,          R\         /# r   r   )r   s   "r   r   r   k  sM     2; 2;2;2; 2; $	2;
 %2; 2;r   c                \   ^ V P                   9   d   \        P                  ! R^R7       V # \        W4      p\	        W14      pV\
        P                  ! V4      ,          p\        P                  ! 4       ;_uu_ 4        V P                  ^ WtR7      uuRRR4       #   + '       g   i     R# ; i)a+  Fill the input `Tensor` with values using a Kaiming normal distribution.

The method is described in `Delving deep into rectifiers: Surpassing
human-level performance on ImageNet classification` - He, K. et al. (2015).
The resulting tensor will have values sampled from
:math:`\mathcal{N}(0, \text{std}^2)` where

.. math::
    \text{std} = \frac{\text{gain}}{\sqrt{\text{fan\_mode}}}

Also known as He initialization.

Args:
    tensor: an n-dimensional `torch.Tensor`
    a: the negative slope of the rectifier used after this layer (only
        used with ``'leaky_relu'``)
    mode: either ``'fan_in'`` (default) or ``'fan_out'``. Choosing ``'fan_in'``
        preserves the magnitude of the variance of the weights in the
        forward pass. Choosing ``'fan_out'`` preserves the magnitudes in the
        backwards pass.
    nonlinearity: the non-linear function (`nn.functional` name),
        recommended to use only with ``'relu'`` or ``'leaky_relu'`` (default).
    generator: the torch Generator to sample from (default: None)

Examples:
    >>> w = torch.empty(3, 5)
    >>> nn.init.kaiming_normal_(w, mode="fan_out", nonlinearity="relu")

Note:
    Be aware that ``fan_in`` and ``fan_out`` are calculated assuming
    that the weight matrix is used in a transposed manner,
    (i.e., ``x @ w.T`` in ``Linear`` layers, where ``w.shape = [fan_out, fan_in]``).
    This is important for correct initialization.
    If you plan to use ``x @ w``, where ``w.shape = [fan_in, fan_out]``,
    pass in a transposed weight matrix, i.e. ``nn.init.kaiming_normal_(w.T, ...)``.
r   r8   r   N)
r   r=   r>   r   r   r3   r5   r   r   r'   )r   r   rR   rh   r   r   r   r%   s   &&&&&   r   kaiming_normal_r   k  ss    V 	FLLDQRS
 
.C,*D
3
C	~~a~: 
s   <BB+	c                j    V ^8  d   QhR\         R\        R\        P                  R,          R\         /# r   r   )r   s   "r   r   r     s9     0 00
0 %0 	0r   c                   V P                  4       ^8  d   \        R4      hV P                  4       ^ 8X  g   V P                  '       d   V # V P	                  ^ 4      pV P                  4       V,          pV P                  W434      P                  ^ ^VR7      pW48  d   VP                  4        \        P                  P                  V4      w  rg\        P                  ! V^ 4      pVP                  4       p	Wi,          pW48  d   VP                  4        \        P                  ! 4       ;_uu_ 4        V P                  V4      P                  V4       V P!                  V4       RRR4       V #   + '       g   i     T # ; i)al  Fill the input `Tensor` with a (semi) orthogonal matrix.

Described in `Exact solutions to the nonlinear dynamics of learning in deep
linear neural networks` - Saxe, A. et al. (2013). The input tensor must have
at least 2 dimensions, and for tensors with more than 2 dimensions the
trailing dimensions are flattened.

Args:
    tensor: an n-dimensional `torch.Tensor`, where :math:`n \geq 2`
    gain: optional scaling factor
    generator: the torch Generator to sample from (default: None)

Examples:
    >>> # xdoctest: +REQUIRES(env:TORCH_DOCTEST_LAPACK)
    >>> w = torch.empty(3, 5)
    >>> nn.init.orthogonal_(w)
z4Only tensors with 2 or more dimensions are supportedr   N)r   r|   numelr<   r   	new_emptyr'   t_r   linalgqrdiagsignr   view_asrD   rJ   )
r   r   r   rowscols	flattenedqrr   phs
   &&&       r   orthogonal_r     s   , QOPP||~fnnn;;q>D<<>T!D  $.66q!y6QI{ <<??9%DA

1aA	
BGA{		q"D 
 M 
 Ms   /2E++E<	c          
      v    V ^8  d   QhR\         R\        R\        R\        P                  R,          R\         /# )r   r   sparsityr%   r   Nr   r   )r   s   "r   r   r     sC     & &&& 
& %	&
 &r   c                   V P                  4       ^8w  d   \        R4      hV P                  '       d   V # V P                  w  rE\        P
                  ! W,          4      p\        P                  ! 4       ;_uu_ 4        V P                  ^ W#R7       \        V4       F$  p\        P                  ! V4      pVRV p	^ W	V3&   K&  	  RRR4       V #   + '       g   i     T # ; i)aZ  Fill the 2D input `Tensor` as a sparse matrix.

The non-zero elements will be drawn from the normal distribution
:math:`\mathcal{N}(0, 0.01)`, as described in `Deep learning via
Hessian-free optimization` - Martens, J. (2010).

Args:
    tensor: an n-dimensional `torch.Tensor`
    sparsity: The fraction of elements in each column to be set to zero
    std: the standard deviation of the normal distribution used to generate
        the non-zero values
    generator: the torch Generator to sample from (default: None)

Examples:
    >>> w = torch.empty(3, 5)
    >>> nn.init.sparse_(w, sparsity=0.1)
r   r   N)r   r|   r<   r   r3   ceilr   r   r'   r   randperm)
r   r   r%   r   r   r   	num_zeroscol_idxrow_indiceszero_indicess
   &&&&      r   sparse_r     s    . aGHH~~~JD		(/*I	q#3T{G...K&z	2L,-F() # 
 M 
 Ms   <ACC	c                t    V ^8  d   QhR\         \        \        3,          R\         \        \        3,          /# )r   methr   )r   r	   r   )r   s   "r   r   r     s,      (2r6* xB/? r   c                 t   a aa S P                   oSR R oR V VV3R llpRS RS RS R2Vn        SVn         V# )Nc                d    V ^8  d   QhR\         P                  R\         P                  R\        /# )r   argskwargsr   )r	   r   r   r   )r   s   "r   r   %_make_deprecate.<locals>.__annotate__  s)     % %rww %")) % %r   c                  \   < \         P                  ! R S RS R2\        ^R7       S! V / VB # )z	`nn.init.z)` is now deprecated in favor of `nn.init.z`.r8   )r=   r>   FutureWarning)r   r   r   new_nameold_names   *,r   deprecated_init(_make_deprecate.<locals>.deprecated_init  s;    z!J8*TVW	

 T$V$$r   z
    z_(...)

    .. warning::
        This method is now deprecated in favor of :func:`torch.nn.init.z"`.

    See :func:`~torch.nn.init.z` for details.)__name____doc__)r   r   r   r   s   f @@r   _make_deprecater     sa    }}H}H% %$J H IQz R'j:O  (Or   )r   r   r'   r   r   r   r   r   r   r   r   r   r   r   r   uniformnormalconstantr   diracxavier_uniformxavier_normalkaiming_uniformkaiming_normal
orthogonalsparse)rs   rt   ru   rv   rw   rx   ry   rm   rn   ro   rp   rr   )r   r   r]   )r   r1   N)r   r1   g       r2   N)   )r1   N)    r   rp   N)r  N)rq   N)3r   r3   r=   collections.abcr   typingr   r   typing_extensionsr   r   r   __all__r   r	   rj   r   r    r)   rW   r`   rf   r   r   r'   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r  r  r  r  r  r   r   r   <module>r     sG   N   $ # '  > T]t_  &':>JZ!

GET66:6!PH'$
'
"*5p*7B9>3>CB2;j0f&T. (
#		!9%d 1/!"23 1[)
		!r   