U
    Ãmœd.  ã                   @   sv   d Z ddlmZ ddlZddlmZ G dd„ dƒZG dd„ dƒZG d	d
„ d
ƒZ	G dd„ dƒZ
dd„ ZG dd„ dƒZdS )a!  
Utilities for cross validation.

taken from scikits.learn

# Author: Alexandre Gramfort <alexandre.gramfort@inria.fr>,
#         Gael Varoquaux    <gael.varoquaux@normalesup.org>
# License: BSD Style.
# $Id$

changes to code by josef-pktd:
 - docstring formatting: underlines of headers

é    )ÚlrangeN)Úcombinationsc                   @   s(   e Zd ZdZdd„ Zdd„ Zdd„ ZdS )	ÚLeaveOneOutzs
    Leave-One-Out cross validation iterator:
    Provides train/test indexes to split data in train test sets
    c                 C   s
   || _ dS )a9  
        Leave-One-Out cross validation iterator:
        Provides train/test indexes to split data in train test sets

        Parameters
        ----------
        n: int
            Total number of elements

        Examples
        --------
        >>> from scikits.learn import cross_val
        >>> X = [[1, 2], [3, 4]]
        >>> y = [1, 2]
        >>> loo = cross_val.LeaveOneOut(2)
        >>> for train_index, test_index in loo:
        ...    print "TRAIN:", train_index, "TEST:", test_index
        ...    X_train, X_test, y_train, y_test = cross_val.split(train_index, test_index, X, y)
        ...    print X_train, X_test, y_train, y_test
        TRAIN: [False  True] TEST: [ True False]
        [[3 4]] [[1 2]] [2] [1]
        TRAIN: [ True False] TEST: [False  True]
        [[1 2]] [[3 4]] [1] [2]
        N)Ún)Úselfr   © r   ú\/home/sam/Atlas/atlas_env/lib/python3.8/site-packages/statsmodels/sandbox/tools/cross_val.pyÚ__init__   s    zLeaveOneOut.__init__c                 c   sB   | j }t|ƒD ].}tj|td�}d||< t |¡}||fV  qd S ©N©ZdtypeT)r   ÚrangeÚnpÚzerosÚboolÚlogical_not)r   r   ÚiÚ
test_indexÚtrain_indexr   r   r   Ú__iter__8   s    
zLeaveOneOut.__iter__c                 C   s   d| j j| j j| jf S ©Nz%s.%s(n=%i)©Ú	__class__Ú
__module__Ú__name__r   ©r   r   r   r   Ú__repr__A   s    þzLeaveOneOut.__repr__N©r   r   Ú__qualname__Ú__doc__r	   r   r   r   r   r   r   r      s   	r   c                   @   s(   e Zd ZdZdd„ Zdd„ Zdd„ ZdS )	Ú	LeavePOutzq
    Leave-P-Out cross validation iterator:
    Provides train/test indexes to split data in train test sets
    c                 C   s   || _ || _dS )aV  
        Leave-P-Out cross validation iterator:
        Provides train/test indexes to split data in train test sets

        Parameters
        ----------
        n: int
            Total number of elements
        p: int
            Size test sets

        Examples
        --------
        >>> from scikits.learn import cross_val
        >>> X = [[1, 2], [3, 4], [5, 6], [7, 8]]
        >>> y = [1, 2, 3, 4]
        >>> lpo = cross_val.LeavePOut(4, 2)
        >>> for train_index, test_index in lpo:
        ...    print "TRAIN:", train_index, "TEST:", test_index
        ...    X_train, X_test, y_train, y_test = cross_val.split(train_index, test_index, X, y)
        TRAIN: [False False  True  True] TEST: [ True  True False False]
        TRAIN: [False  True False  True] TEST: [ True False  True False]
        TRAIN: [False  True  True False] TEST: [ True False False  True]
        TRAIN: [ True False False  True] TEST: [False  True  True False]
        TRAIN: [ True False  True False] TEST: [False  True False  True]
        TRAIN: [ True  True False False] TEST: [False False  True  True]
        N)r   Úp)r   r   r    r   r   r   r	   P   s    zLeavePOut.__init__c                 c   sX   | j }| j}tt|ƒ|ƒ}|D ]4}tj|td�}d|t |¡< t |¡}||fV  qd S r
   )	r   r    r   r   r   r   r   Úarrayr   )r   r   r    ÚcombÚidxr   r   r   r   r   r   p   s    
zLeavePOut.__iter__c                 C   s   d| j j| j j| j| jf S )Nz%s.%s(n=%i, p=%i))r   r   r   r   r    r   r   r   r   r   {   s    üzLeavePOut.__repr__Nr   r   r   r   r   r   J   s    r   c                   @   s(   e Zd ZdZdd„ Zdd„ Zdd„ ZdS )	ÚKFoldzm
    K-Folds cross validation iterator:
    Provides train/test indexes to split data in train test sets
    c                 C   s@   |dkst tdƒƒ‚||k s0t td||f ƒƒ‚|| _|| _dS )a—  
        K-Folds cross validation iterator:
        Provides train/test indexes to split data in train test sets

        Parameters
        ----------
        n: int
            Total number of elements
        k: int
            number of folds

        Examples
        --------
        >>> from scikits.learn import cross_val
        >>> X = [[1, 2], [3, 4], [1, 2], [3, 4]]
        >>> y = [1, 2, 3, 4]
        >>> kf = cross_val.KFold(4, k=2)
        >>> for train_index, test_index in kf:
        ...    print "TRAIN:", train_index, "TEST:", test_index
        ...    X_train, X_test, y_train, y_test = cross_val.split(train_index, test_index, X, y)
        TRAIN: [False False  True  True] TEST: [ True  True False False]
        TRAIN: [ True  True False False] TEST: [False False  True  True]

        Notes
        -----
        All the folds have size trunc(n/k), the last one has the complementary
        r   zcannot have k below 1z cannot have k=%d greater than %dN)ÚAssertionErrorÚ
ValueErrorr   Úk)r   r   r'   r   r   r   r	   ‹   s    zKFold.__init__c                 c   sˆ   | j }| j}tt || ¡ƒ}t|ƒD ]\}tj|td�}||d k r^d||| |d | …< nd||| d …< t |¡}||fV  q&d S )Nr   é   T)	r   r'   Úintr   Úceilr   r   r   r   )r   r   r'   Újr   r   r   r   r   r   r   ­   s    
zKFold.__iter__c                 C   s   d| j j| j j| j| jf S )Nz%s.%s(n=%i, k=%i))r   r   r   r   r'   r   r   r   r   r   ¼   s    üzKFold.__repr__Nr   r   r   r   r   r$   …   s   "r$   c                   @   s(   e Zd ZdZdd„ Zdd„ Zdd„ ZdS )	ÚLeaveOneLabelOutzy
    Leave-One-Label_Out cross-validation iterator:
    Provides train/test indexes to split data in train test sets
    c                 C   s
   || _ dS )aõ  
        Leave-One-Label_Out cross validation:
        Provides train/test indexes to split data in train test sets

        Parameters
        ----------
        labels : list
                List of labels

        Examples
        --------
        >>> from scikits.learn import cross_val
        >>> X = [[1, 2], [3, 4], [5, 6], [7, 8]]
        >>> y = [1, 2, 1, 2]
        >>> labels = [1, 1, 2, 2]
        >>> lol = cross_val.LeaveOneLabelOut(labels)
        >>> for train_index, test_index in lol:
        ...    print "TRAIN:", train_index, "TEST:", test_index
        ...    X_train, X_test, y_train, y_test = cross_val.split(train_index,             test_index, X, y)
        ...    print X_train, X_test, y_train, y_test
        TRAIN: [False False  True  True] TEST: [ True  True False False]
        [[5 6]
        [7 8]] [[1 2]
        [3 4]] [1 2] [1 2]
        TRAIN: [ True  True False False] TEST: [False False  True  True]
        [[1 2]
        [3 4]] [[5 6]
        [7 8]] [1 2] [1 2]
        N)Úlabels)r   r-   r   r   r   r	   Ì   s    zLeaveOneLabelOut.__init__c                 c   sV   t j| jdd�}t  |¡D ]6}t jt|ƒtd�}d|||k< t  |¡}||fV  qd S )NT)Úcopyr   )r   r!   r-   Úuniquer   Úlenr   r   )r   r-   r   r   r   r   r   r   r   î   s    
zLeaveOneLabelOut.__iter__c                 C   s   d| j j| j j| jf S )Nz%s.%s(labels=%s))r   r   r   r-   r   r   r   r   r   ø   s
    ýzLeaveOneLabelOut.__repr__Nr   r   r   r   r   r,   Æ   s   "
r,   c                 G   s@   g }|D ]2}t  |¡}||  }|| }| |¡ | |¡ q|S )zx
    For each arg return a train and test subsets defined by indexes provided
    in train_indexes and test_indexes
    )r   Z
asanyarrayÚappend)Ztrain_indexesZtest_indexesÚargsÚretÚargZ	arg_trainZarg_testr   r   r   Úsplit   s    

r5   c                   @   s*   e Zd ZdZddd„Zdd„ Zd	d
„ ZdS )Ú
KStepAheadzn
    KStepAhead cross validation iterator:
    Provides fit/test indexes to split data in sequential sets
    r(   NTc                 C   s<   || _ || _|dkr&tt |d ¡ƒ}|| _|| _|| _dS )a=  
        KStepAhead cross validation iterator:
        Provides train/test indexes to split data in train test sets

        Parameters
        ----------
        n: int
            Total number of elements
        k : int
            number of steps ahead
        start : int
            initial size of data for fitting
        kall : bool
            if true. all values for up to k-step ahead are included in the test index.
            If false, then only the k-th step ahead value is returnd


        Notes
        -----
        I do not think this is really useful, because it can be done with
        a very simple loop instead.
        Useful as a plugin, but it could return slices instead for faster array access.

        Examples
        --------
        >>> from scikits.learn import cross_val
        >>> X = [[1, 2], [3, 4]]
        >>> y = [1, 2]
        >>> loo = cross_val.LeaveOneOut(2)
        >>> for train_index, test_index in loo:
        ...    print "TRAIN:", train_index, "TEST:", test_index
        ...    X_train, X_test, y_train, y_test = cross_val.split(train_index, test_index, X, y)
        ...    print X_train, X_test, y_train, y_test
        TRAIN: [False  True] TEST: [ True False]
        [[3 4]] [[1 2]] [2] [1]
        TRAIN: [ True False] TEST: [False  True]
        [[1 2]] [[3 4]] [1] [2]
        Ng      Ð?)r   r'   r)   r   ÚtruncÚstartÚkallÚreturn_slice)r   r   r'   r8   r9   r:   r   r   r   r	      s    'zKStepAhead.__init__c           	      c   sê   | j }| j}| j}| jrpt||| ƒD ]F}td |d ƒ}| jrLt||| ƒ}nt|| d || ƒ}||fV  q&nvt||| ƒD ]f}tj|t	d�}d|d |…< tj|t	d�}| jrÂd|||| …< nd||| d || …< ||fV  q~d S )Nr(   r   T)
r   r'   r8   r:   r   Úslicer9   r   r   r   )	r   r   r'   r8   r   Ztrain_sliceZ
test_slicer   r   r   r   r   r   P  s$    zKStepAhead.__iter__c                 C   s   d| j j| j j| jf S r   r   r   r   r   r   r   k  s    þzKStepAhead.__repr__)r(   NTTr   r   r   r   r   r6     s   
0r6   )r   Zstatsmodels.compat.pythonr   Únumpyr   Ú	itertoolsr   r   r   r$   r,   r5   r6   r   r   r   r   Ú<module>   s   4;A: