Ë
    ÜÍ:j?  ã                   ó  — d Z ddlmZmZ ddlmZ ddlmZ ddlm	Z	 ddl
mZmZmZ ddlmZ e G d„ d	«      «       Zd
„ Z G d„ de«      Z G d„ de«      Z G d„ de«      Z G d„ de«      Z G d„ de«      Z G d„ de«      ZeeeeedœZy)zM
Module contains classes for invertible (and differentiable) link functions.
é    )ÚABCÚabstractmethod)Ú	dataclass)Úulp)Úgmean)Ú_expitÚ_logitÚget_namespace©Úsoftmaxc                   óB   — e Zd ZU eed<   eed<   eed<   eed<   d„ Zd„ Zy)ÚIntervalÚlowÚhighÚlow_inclusiveÚhigh_inclusivec                 ó‚   — | j                   | j                  kD  r&t        d| j                   › d| j                  › d�«      ‚y)zCheck that low <= highz#One must have low <= high; got low=z, high=ú.N)r   r   Ú
ValueError)Úselfs    úg/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/sklearn/_loss/link.pyÚ__post_init__zInterval.__post_init__   s?   € à�8‰8�d—i‘iÒÜØ5°d·h±h°Z¸wÀtÇyÁyÀkÐQRÐSóð ð  ó    c                 óŠ  — t        |«      \  }}| j                  r|j                  || j                  «      }n|j	                  || j                  «      }|j                  |«      sy| j                  r|j                  || j                  «      }n|j                  || j                  «      }t        |j                  |«      «      S )zóTest whether all values of x are in interval range.

        Parameters
        ----------
        x : ndarray
            Array whose elements are tested to be in interval range.

        Returns
        -------
        result : bool
        F)r
   r   Úgreater_equalr   ÚgreaterÚallr   Ú
less_equalr   ÚlessÚbool)r   ÚxÚxpÚ_r   r   s         r   ÚincludeszInterval.includes    s™   € ô ˜aÓ ‰ˆˆAØ×ÒØ×"Ñ" 1 d§h¡hÓ/‰Cà—*‘*˜Q §¡Ó)ˆCà�v‰v�cŒ{Øà×ÒØ—=‘=  D§I¡IÓ.‰Dà—7‘7˜1˜dŸi™iÓ(ˆDô �B—F‘F˜4“LÓ!Ð!r   N)Ú__name__Ú
__module__Ú__qualname__ÚfloatÚ__annotations__r    r   r$   © r   r   r   r      s"   … à	ƒJØ
ƒKØÓØÓòó"r   r   c                 ó   — dt        d«      z  }| j                  t        d«       k(  rd}n:| j                  dk  r| j                  d|z
  z  |z   }n| j                  d|z   z  |z   }| j                  t        d«      k(  rd}n:| j                  dk  r| j                  d|z   z  |z
  }n| j                  d|z
  z  |z
  }t        |«      t        |«      fS )zÞGenerate values low and high to be within the interval range.

    This is used in tests only.

    Returns
    -------
    low, high : tuple of floats
        The returned values low and high lie within the interval.
    é
   é   Úinfg    _ Âr   g    _ B)r   r   r(   r   )ÚintervalÚepsr   r   s       r   Ú_inclusive_low_highr1   >   sÄ   € ð Œs�1‹v‰+€CØ‡|�|œ˜e›�}Ò$Ø‰Ø	�‰˜Ò	Ø�l‰l˜a #™gÑ&¨Ñ,‰à�l‰l˜a #™gÑ&¨Ñ,ˆà‡}�}œ˜e›Ò$Ø‰Ø	�‰˜Ò	Ø�}‰}  C¡Ñ(¨3Ñ.‰à�}‰}  C¡Ñ(¨3Ñ.ˆä�‹:”u˜T“{Ð"Ð"r   c                   ód   — e Zd ZdZdZ e ed«        ed«      dd«      Zed„ «       Z	ed„ «       Z
y)ÚBaseLinka   Abstract base class for differentiable, invertible link functions.

    Convention:
        - link function g: raw_prediction = g(y_pred)
        - inverse link h: y_pred = h(raw_prediction)

    For (generalized) linear models, `raw_prediction = X @ coef` is the so
    called linear predictor, and `y_pred = h(raw_prediction)` is the predicted
    conditional (on X) expected value of the target `y_true`.

    The methods are not implemented as staticmethods in case a link function needs
    parameters.
    Fr.   c                  ó   — y)an  Compute the link function g(y_pred).

        The link function maps (predicted) target values to raw predictions,
        i.e. `g(y_pred) = raw_prediction`.

        Parameters
        ----------
        y_pred : array
            Predicted target values.

        Returns
        -------
        array
            Output array, element-wise link function.
        Nr*   ©r   Úy_preds     r   ÚlinkzBaseLink.linkp   ó   � r   c                  ó   — y)a¢  Compute the inverse link function h(raw_prediction).

        The inverse link function maps raw predictions to predicted target
        values, i.e. `h(raw_prediction) = y_pred`.

        Parameters
        ----------
        raw_prediction : array
            Raw prediction values (in link space).

        Returns
        -------
        array
            Output array, element-wise inverse link function.
        Nr*   ©r   Úraw_predictions     r   ÚinversezBaseLink.inverse‚   r8   r   N)r%   r&   r'   Ú__doc__Úis_multiclassr   r(   Úinterval_y_predr   r7   r<   r*   r   r   r3   r3   Z   sP   „ ñð €Mñ
 ¡ e£˜}©e°E«l¸EÀ5ÓI€Oàñó ðð" ñó ñr   r3   c                   ó   — e Zd ZdZd„ ZeZy)ÚIdentityLinkz"The identity link function g(x)=x.c                 ó   — |S ©Nr*   r5   s     r   r7   zIdentityLink.link˜   s   € Øˆr   N)r%   r&   r'   r=   r7   r<   r*   r   r   rA   rA   •   s   „ Ù,òð �Gr   rA   c                   ó>   — e Zd ZdZ ed ed«      dd«      Zd„ Zd„ Zy)ÚLogLinkz"The log link function g(x)=log(x).r   r.   Fc                 ó@   — t        |«      \  }}|j                  |«      S rC   )r
   Úlog)r   r6   r"   r#   s       r   r7   zLogLink.link£   s   € Ü˜fÓ%‰ˆˆAØ�v‰v�f‹~Ðr   c                 ó@   — t        |«      \  }}|j                  |«      S rC   )r
   Úexp©r   r;   r"   r#   s       r   r<   zLogLink.inverse§   s   € Ü˜nÓ-‰ˆˆAØ�v‰v�nÓ%Ð%r   N)	r%   r&   r'   r=   r   r(   r?   r7   r<   r*   r   r   rE   rE   ž   s#   „ Ù,á˜q¡%¨£,°°uÓ=€Oòó&r   rE   c                   ó2   — e Zd ZdZ edddd«      Zd„ Zd„ Zy)Ú	LogitLinkz&The logit link function g(x)=logit(x).r   r-   Fc                 ó   — t        |«      S rC   ©r	   r5   s     r   r7   zLogitLink.link±   s   € Ü�f‹~Ðr   c                 ó   — t        |«      S rC   ©r   r:   s     r   r<   zLogitLink.inverse´   s   € Ü�nÓ%Ð%r   N©r%   r&   r'   r=   r   r?   r7   r<   r*   r   r   rL   rL   ¬   s   „ Ù0á˜q ! U¨EÓ2€Oòó&r   rL   c                   ó2   — e Zd ZdZ edddd«      Zd„ Zd„ Zy)ÚHalfLogitLinkzZHalf the logit link function g(x)=1/2 * logit(x).

    Used for the exponential loss.
    r   r-   Fc                 ó   — dt        |«      z  S )Ng      à?rN   r5   s     r   r7   zHalfLogitLink.linkÀ   s   € Ø”V˜F“^Ñ#Ð#r   c                 ó   — t        d|z  «      S )Né   rP   r:   s     r   r<   zHalfLogitLink.inverseÃ   s   € Ü�a˜.Ñ(Ó)Ð)r   NrQ   r*   r   r   rS   rS   ¸   s#   „ ññ
 ˜q ! U¨EÓ2€Oò$ó*r   rS   c                   ó<   — e Zd ZdZdZ edddd«      Zd„ Zd„ Zd„ Z	y	)
ÚMultinomialLogitaš  The symmetric multinomial logit function.

    Convention:
        - y_pred.shape = raw_prediction.shape = (n_samples, n_classes)

    Notes:
        - The inverse link h is the softmax function.
        - The sum is over the second axis, i.e. axis=1 (n_classes).

    We have to choose additional constraints in order to make

        y_pred[k] = exp(raw_pred[k]) / sum(exp(raw_pred[k]), k=0..n_classes-1)

    for n_classes classes identifiable and invertible.
    We choose the symmetric side constraint where the geometric mean response
    is set as reference category, see [2]:

    The symmetric multinomial logit link function for a single data point is
    then defined as

        raw_prediction[k] = g(y_pred[k]) = log(y_pred[k]/gmean(y_pred))
        = log(y_pred[k]) - mean(log(y_pred)).

    Note that this is equivalent to the definition in [1] and implies mean
    centered raw predictions:

        sum(raw_prediction[k], k=0..n_classes-1) = 0.

    For linear models with raw_prediction = X @ coef, this corresponds to
    sum(coef[k], k=0..n_classes-1) = 0, i.e. the sum over classes for every
    feature is zero.

    Reference
    ---------
    .. [1] Friedman, Jerome; Hastie, Trevor; Tibshirani, Robert. "Additive
        logistic regression: a statistical view of boosting" Ann. Statist.
        28 (2000), no. 2, 337--407. doi:10.1214/aos/1016218223.
        https://projecteuclid.org/euclid.aos/1016218223

    .. [2] Zahid, Faisal Maqbool and Gerhard Tutz. "Ridge estimation for
        multinomial logit models with symmetric side constraints."
        Computational Statistics 28 (2013): 1017-1034.
        http://epub.ub.uni-muenchen.de/11001/1/tr067.pdf
    Tr   r-   Fc                 óX   — t        |«      \  }}||j                  |d¬«      d d …d f   z
  S ©Nr-   )Úaxis)r
   ÚmeanrJ   s       r   Úsymmetrize_raw_predictionz*MultinomialLogit.symmetrize_raw_predictionø   s1   € Ü˜nÓ-‰ˆˆAØ §¡¨¸Q Ó ?ÂÀ4ÀÑ HÑHÐHr   c                 ón   — t        |«      \  }}t        |d¬«      }|j                  ||d d …d f   z  «      S rZ   )r
   r   rG   )r   r6   r"   r#   Úgms        r   r7   zMultinomialLogit.linkü   s8   € Ü˜fÓ%‰ˆˆAä�6 Ô"ˆØ�v‰v�f˜r¢! T '™{Ñ*Ó+Ð+r   c                 ó   — t        |«      S rC   r   r:   s     r   r<   zMultinomialLogit.inverse  s   € Ü�~Ó&Ð&r   N)
r%   r&   r'   r=   r>   r   r?   r]   r7   r<   r*   r   r   rX   rX   Ç   s/   „ ñ+ðZ €MÙ˜q ! U¨EÓ2€OòIò,ó'r   rX   )ÚidentityrG   ÚlogitÚ
half_logitÚmultinomial_logitN)r=   Úabcr   r   Údataclassesr   Úmathr   Úscipy.statsr   Úsklearn.utils._array_apir   r	   r
   Úsklearn.utils.extmathr   r   r1   r3   rA   rE   rL   rS   rX   Ú_LINKSr*   r   r   ú<module>rl      s¢   ðñ÷ $Ý !Ý å ç BÑ BÝ )ð ÷("ð ("ó ð("òV#ô88ˆsô 8ôv�8ô ô&ˆhô &ô	&�ô 	&ô*�Hô *ô<'�xô <'ð@ ØØØØ)ñ
�r   