U
    ÃmœdH  ã                   @   s:   d Z ddlZddlmZ ddlmZmZ G dd„ dƒZdS )zL
Created on Sun May 10 08:23:48 2015

Author: Josef Perktold
License: BSD-3
é    Né   )ÚNonePenalty)Úapprox_fprime_csÚapprox_fprimec                       s’   e Zd ZdZ‡ fdd„Zddd„Zd‡ fdd„	Zd‡ fd	d
„	Zddd„Zd‡ fdd„	Z	d‡ fdd„	Z
ddd„Zd‡ fdd„	Zd ‡ fdd„	Z‡  ZS )!ÚPenalizedMixina  Mixin class for Maximum Penalized Likelihood

    Parameters
    ----------
    args and kwds for the model super class
    penal : None or instance of Penalized function class
        If penal is None, then NonePenalty is used.
    pen_weight : float or None
        factor for weighting the penalization term.
        If None, then pen_weight is set to nobs.


    TODO: missing **kwds or explicit keywords

    TODO: do we adjust the inherited docstrings?
    We would need templating to add the penalization parameters
    c                    sŽ   |  dd ¡| _|  dd ¡| _tt| ƒj||Ž | jd krDt| jƒ| _| jd kr\tƒ | _d| _| j	 
ddg¡ t| dg ƒ| _| j 
ddg¡ d S )NÚpenalÚ
pen_weightr   Ú_null_drop_keys)Úpopr   r   Úsuperr   Ú__init__ÚlenZendogr   Z
_init_keysÚextendÚgetattrr	   )ÚselfÚargsÚkwds©Ú	__class__© úT/home/sam/Atlas/atlas_env/lib/python3.8/site-packages/statsmodels/base/_penalized.pyr   !   s    

zPenalizedMixin.__init__Nc                 K   s0   |d kr,t | dƒr(|  |¡}|  |¡}nd}|S )NZ	scaletyper   )ÚhasattrZpredictZestimate_scale)r   ÚparamsÚscaler   Úmur   r   r   Ú_handle_scale8   s    

zPenalizedMixin._handle_scalec                    sX   |dkr| j }tt| ƒj|f|Ž}|dkrT| j|f|Ž}|d| | | j |¡ 8 }|S )z3
        Log-likelihood of model at params
        Nr   r   )r   r   r   Úlogliker   r   Úfunc)r   r   r   r   Úllfr   r   r   r   r   D   s    zPenalizedMixin.loglikec                    sj   |dkr| j }tt| ƒj|f|Ž}t|jd ƒ}|dkrf| j|f|Ž}|d| | | | j |¡ 8 }|S )z@
        Log-likelihood of model observations at params
        Nr   r   )	r   r   r   Ú
loglikeobsÚfloatÚshaper   r   r   )r   r   r   r   r   Znobs_llfr   r   r   r   r   R   s     zPenalizedMixin.loglikeobsÚfdc                    sR   ˆdkrˆj ‰‡ ‡‡fdd„}|dkr0t||ƒS |dkrFt||dd�S tdƒ‚dS )	z4score based on finite difference derivative
        Nc                    s   ˆj | fdˆiˆ —ŽS ©Nr   ©r   ©Úp©r   r   r   r   r   Ú<lambda>h   ó    z.PenalizedMixin.score_numdiff.<locals>.<lambda>Úcsr"   T)Zcenteredz-method not recognized, should be "fd" or "cs")r   r   r   Ú
ValueError)r   r   r   Úmethodr   r   r   r'   r   Úscore_numdiffb   s    
zPenalizedMixin.score_numdiffc                    sX   |dkr| j }tt| ƒj|f|Ž}|dkrT| j|f|Ž}|d| | | j |¡ 8 }|S )z-
        Gradient of model at params
        Nr   r   )r   r   r   Úscorer   r   Úderiv)r   r   r   r   Úscr   r   r   r   r.   q   s    zPenalizedMixin.scorec                    sj   |dkr| j }tt| ƒj|f|Ž}t|jd ƒ}|dkrf| j|f|Ž}|d| | | | j |¡ 8 }|S )z:
        Gradient of model observations at params
        Nr   r   )	r   r   r   Ú	score_obsr    r!   r   r   r/   )r   r   r   r   r0   Znobs_scr   r   r   r   r1      s     zPenalizedMixin.score_obsc                    s4   ˆdkrˆj ‰‡ ‡‡fdd„}ddlm} |||ƒS )z6hessian based on finite difference derivative
        Nc                    s   ˆj | fdˆiˆ —ŽS r#   r$   r%   r'   r   r   r(   “   r)   z0PenalizedMixin.hessian_numdiff.<locals>.<lambda>r   )Úapprox_hess)r   Ústatsmodels.tools.numdiffr2   )r   r   r   r   r   r2   r   r'   r   Úhessian_numdiffŽ   s
    zPenalizedMixin.hessian_numdiffc                    s‚   |dkr| j }tt| ƒj|f|Ž}|dkr~| j|f|Ž}| j |¡}|jdkrj|d| t 	|| ¡ 8 }n|d| | | 8 }|S )z,
        Hessian of model at params
        Nr   r   )
r   r   r   Úhessianr   r   Zderiv2ÚndimÚnpZdiag)r   r   r   r   Zhessr   Úhr   r   r   r5   ˜   s    
zPenalizedMixin.hessianc           
         sÔ   ddl m} ddlm} t| ||fƒr4| ddi¡ |dkr@d}|dkrLd}tt| ƒjf d|i|—Ž}|dkrr|S |d	kr~d
}t	 
t	 |j¡|k ¡d }t	 
t	 |j¡|k¡d }| ¡ rÌ| j|f|Ž}	|	S |S dS )aƒ  minimize negative penalized log-likelihood

        Parameters
        ----------
        method : None or str
            Method specifies the scipy optimizer as in nonlinear MLE models.
        trim : {bool, float}
            Default is False or None, which uses no trimming.
            If trim is True or a float, then small parameters are set to zero.
            If True, then a default threshold is used. If trim is a float, then
            it will be used as threshold.
            The default threshold is currently 1e-4, but it will change in
            future and become penalty function dependent.
        kwds : extra keyword arguments
            This keyword arguments are treated in the same way as in the
            fit method of the underlying model class.
            Specifically, additional optimizer keywords and cov_type related
            keywords can be added.
        r   )ÚGLMGam)ÚGLMZmax_start_irlsNZbfgsFr,   Tg-Cëâ6?)Z*statsmodels.gam.generalized_additive_modelr9   Z+statsmodels.genmod.generalized_linear_modelr:   Ú
isinstanceÚupdater   r   Úfitr7   ZnonzeroÚabsr   ÚanyZ
_fit_zeros)
r   r,   Ztrimr   r9   r:   ÚresZ
drop_indexZ
keep_indexZres_auxr   r   r   r=   ª   s&    zPenalizedMixin.fit)N)N)N)Nr"   )N)N)N)N)NN)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   r   r   r   r-   r.   r1   r4   r5   r=   Ú__classcell__r   r   r   r   r      s   



r   )	rD   Únumpyr7   Z
_penaltiesr   r3   r   r   r   r   r   r   r   Ú<module>   s   