U
    @¼|e÷  ã                   @   sÒ   d Z ddlmZmZ ddlmZ ddlZddlm	Z	m
Z
 ddlmZ ddlmZ eG d	d
„ d
ƒƒZejfdd„ZG dd„ deƒZG dd„ deƒZG dd„ deƒZG dd„ deƒZG dd„ deƒZeeeedœZdS )zM
Module contains classes for invertible (and differentiable) link functions.
é    )ÚABCÚabstractmethod)Ú	dataclassN)ÚexpitÚlogit)Úgmeané   )Úsoftmaxc                   @   s>   e Zd ZU eed< eed< eed< eed< dd„ Zdd„ Zd	S )
ÚIntervalÚlowÚhighÚlow_inclusiveÚhigh_inclusivec                 C   s*   | j | jkr&td| j › d| j› d�ƒ‚dS )zCheck that low <= highz#One must have low <= high; got low=z, high=Ú.N)r   r   Ú
ValueError)Úself© r   úO/var/www/website-v5/atlas_env/lib/python3.8/site-packages/sklearn/_loss/link.pyÚ__post_init__   s    ÿzInterval.__post_init__c                 C   sd   | j rt || j¡}nt || j¡}t |¡s2dS | jrHt || j¡}nt 	|| j¡}t
t |¡ƒ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   ÚnpÚgreater_equalr   ÚgreaterÚallr   Ú
less_equalr   ÚlessÚbool)r   Úxr   r   r   r   r   Úincludes   s    
zInterval.includesN)Ú__name__Ú
__module__Ú__qualname__ÚfloatÚ__annotations__r   r   r   r   r   r   r   r
      s   
r
   c                 C   sž   dt  |¡j }| jt j kr$d}n0| jdk rB| jd|  | }n| jd|  | }| jt jkrfd}n0| jdk r„| jd|  | }n| jd|  | }||fS )zÔGenerate values low and high to be within the interval range.

    This is used in tests only.

    Returns
    -------
    low, high : tuple
        The returned values low and high lie within the interval.
    é
   g    _ Âr   é   g    _ B)r   ÚfinfoÚepsr   Úinfr   )ÚintervalÚdtyper&   r   r   r   r   r   Ú_inclusive_low_high:   s    


r*   c                   @   sD   e Zd ZdZdZeej ejddƒZe	ddd„ƒZ
e	d	dd„ƒZdS )
Ú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.
    FNc                 C   s   dS )aX  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.
        out : array
            A location into which the result is stored. If provided, it must
            have a shape that the inputs broadcast to. If not provided or None,
            a freshly-allocated array is returned.

        Returns
        -------
        out : array
            Output array, element-wise link function.
        Nr   ©r   Úy_predÚoutr   r   r   Úlinkl   s    zBaseLink.linkc                 C   s   dS )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).
        out : array
            A location into which the result is stored. If provided, it must
            have a shape that the inputs broadcast to. If not provided or None,
            a freshly-allocated array is returned.

        Returns
        -------
        out : array
            Output array, element-wise inverse link function.
        Nr   ©r   Úraw_predictionr.   r   r   r   Úinverse‚   s    zBaseLink.inverse)N)N)r   r   r    Ú__doc__Úis_multiclassr
   r   r'   Úinterval_y_predr   r/   r2   r   r   r   r   r+   V   s   r+   c                   @   s   e Zd ZdZddd„ZeZdS )ÚIdentityLinkz"The identity link function g(x)=x.Nc                 C   s    |d k	rt  ||¡ |S |S d S )N)r   Úcopytor,   r   r   r   r/   œ   s    zIdentityLink.link)N)r   r   r    r3   r/   r2   r   r   r   r   r6   ™   s   
r6   c                   @   s4   e Zd ZdZedejddƒZd	dd„Zd
dd„Z	dS )ÚLogLinkz"The log link function g(x)=log(x).r   FNc                 C   s   t j||d�S ©N©r.   )r   Úlogr,   r   r   r   r/   «   s    zLogLink.linkc                 C   s   t j||d�S r9   )r   Úexpr0   r   r   r   r2   ®   s    zLogLink.inverse)N)N)
r   r   r    r3   r
   r   r'   r5   r/   r2   r   r   r   r   r8   ¦   s   
r8   c                   @   s2   e Zd ZdZeddddƒZd
dd„Zddd	„ZdS )Ú	LogitLinkz&The logit link function g(x)=logit(x).r   r$   FNc                 C   s   t ||d�S r9   )r   r,   r   r   r   r/   ·   s    zLogitLink.linkc                 C   s   t ||d�S r9   )r   r0   r   r   r   r2   º   s    zLogitLink.inverse)N)N)r   r   r    r3   r
   r5   r/   r2   r   r   r   r   r=   ²   s   
r=   c                   @   s>   e Zd ZdZdZeddddƒZdd„ Zdd	d
„Zddd„Z	dS )Ú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                 C   s    |t j|dd�d d …t jf  S )Nr$   ©Úaxis)r   ÚmeanÚnewaxis)r   r1   r   r   r   Úsymmetrize_raw_predictionï   s    z*MultinomialLogit.symmetrize_raw_predictionNc                 C   s,   t |dd�}tj||d d …tjf  |d�S )Nr$   r?   r:   )r   r   r;   rB   )r   r-   r.   Úgmr   r   r   r/   ò   s    zMultinomialLogit.linkc                 C   s4   |d krt |dd�S t ||¡ t |dd� |S d S )NT)ÚcopyF)r	   r   r7   r0   r   r   r   r2   ÷   s
    zMultinomialLogit.inverse)N)N)
r   r   r    r3   r4   r
   r5   rC   r/   r2   r   r   r   r   r>   ¾   s   -
r>   )Úidentityr;   r   Zmultinomial_logit)r3   Úabcr   r   Údataclassesr   Únumpyr   Úscipy.specialr   r   Úscipy.statsr   Úutils.extmathr	   r
   Úfloat64r*   r+   r6   r8   r=   r>   Z_LINKSr   r   r   r   Ú<module>   s&   *CCü