Ë
    5¾ªjî  ã                   ó|   — d Z ddlmZmZ ddlZddlZddlm	Z	 ddl
mZ ddlmZ ddlmZ ddlmZ  G d	„ d
ee	«      Zy)zwClassifier utilities for MySQL Connector/Python.

Provides a scikit-learn compatible classifier backed by HeatWave ML.
é    )ÚOptionalÚUnionN)ÚClassifierMixin)ÚMyBaseMLModel)ÚML_TASK)Ú	copy_dict)ÚMySQLConnectionAbstractc                   óP  — e Zd ZdZ	 	 	 	 ddedee   dee   dee   dee   f
d„Zd	e	e
j                  ej                  f   d
ej                  fd„Zd	e	e
j                  ej                  f   d
ej                  fd„Zd	e	e
j                  ej                  f   d
e
j                  fd„Zy)ÚMyClassifiera  
    MySQL HeatWave scikit-learn compatible classifier estimator.

    Provides prediction and probability output from a model deployed in MySQL,
    and manages fit, explain, and prediction options as per HeatWave ML interface.

    Attributes:
        predict_extra_options (dict): Dictionary of optional parameters passed through
            to the MySQL backend for prediction and probability inference.
        _model (MyModel): Underlying interface for database model operations.
        fit_extra_options (dict): See MyBaseMLModel.

    Args:
        db_connection (MySQLConnectionAbstract): Active MySQL connector DB connection.
        model_name (str, optional): Custom name for the model.
        fit_extra_options (dict, optional): Extra options for fitting.
        explain_extra_options (dict, optional): Extra options for explanations.
        predict_extra_options (dict, optional): Extra options for predict/predict_proba.

    Methods:
        predict(X): Predict class labels.
        predict_proba(X): Predict class probabilities.
    NÚdb_connectionÚ
model_nameÚfit_extra_optionsÚexplain_extra_optionsÚpredict_extra_optionsc                 ó”   — t        j                  | |t        j                  ||¬«       t	        |«      | _        t	        |«      | _        y)a  
        Initialize a MyClassifier.

        Args:
            db_connection: Active MySQL connector database connection.
            model_name: Optional, custom model name.
            fit_extra_options: Optional fit options.
            explain_extra_options: Optional explain options.
            predict_extra_options: Optional predict/predict_proba options.

        Raises:
            DatabaseError:
                If a database connection issue occurs.
                If an operational error occurs during execution.
        )r   r   N)r   Ú__init__r   ÚCLASSIFICATIONr   r   r   )Úselfr   r   r   r   r   s         úU/var/www/html/serviGia/entorno/lib/python3.12/site-packages/mysql/ai/ml/classifier.pyr   zMyClassifier.__init__G   sE   € ô. 	×ÑØØÜ×"Ñ"Ø!Ø/õ	
ô &/Ð/DÓ%EˆÔ"Ü%.Ð/DÓ%EˆÕ"ó    ÚXÚreturnc                 óv   — | j                   j                  || j                  ¬«      }|d   j                  «       S )a  
        Predict class labels for the input features using the MySQL model.

        References:
            https://dev.mysql.com/doc/heatwave/en/mys-hwaml-ml-predict-table.html
                A full list of supported options can be found under "ML_PREDICT_TABLE Options"

        Args:
            X: Input samples as a numpy array or pandas DataFrame.

        Returns:
            ndarray: Array of predicted class labels, shape (n_samples,).

        Raises:
            DatabaseError:
                If provided options are invalid or unsupported,
                or if the model is not initialized, i.e., fit or import has not
                been called
                If a database connection issue occurs.
                If an operational error occurs during execution.
        ©ÚoptionsÚ
Prediction)Ú_modelÚpredictr   Úto_numpy)r   r   Úresults      r   r   zMyClassifier.predicth   s7   € ð0 —‘×$Ñ$ Q°×0JÑ0JÐ$ÓKˆØ�lÑ#×,Ñ,Ó.Ð.r   c                 óþ   ‡— | j                   j                  || j                  ¬«      }t        |d   j                  d   d   j                  «       «      Št        j                  |d   j                  ˆfd„«      «      S )a*  
        Predict class probabilities for the input features using the MySQL model.

        References:
            https://dev.mysql.com/doc/heatwave/en/mys-hwaml-ml-predict-table.html
                A full list of supported options can be found under "ML_PREDICT_TABLE Options"

        Args:
            X: Input samples as a numpy array or pandas DataFrame.

        Returns:
            ndarray: Array of shape (n_samples, n_classes) with class probabilities.

        Raises:
            DatabaseError:
                If provided options are invalid or unsupported,
                or if the model is not initialized, i.e., fit or import has not
                been called
                If a database connection issue occurs.
                If an operational error occurs during execution.
        r   Ú
ml_resultsr   Úprobabilitiesc                 ó:   •— ‰D �cg c]
  }| d   |   ‘Œ c}S c c}w )Nr#   © )Ú	ml_resultÚ
class_nameÚclassess     €r   ú<lambda>z,MyClassifier.predict_proba.<locals>.<lambda>¡   s'   ø€ ØMTö#Ø?I�I˜oÑ.¨zÓ:ò#€ ùò #s   †)	r   r   r   ÚsortedÚilocÚkeysÚnpÚstackÚmap)r   r   r    r(   s      @r   Úpredict_probazMyClassifier.predict_probaƒ   su   ø€ ð0 —‘×$Ñ$ Q°×0JÑ0JÐ$ÓKˆä˜ Ñ-×2Ñ2°1Ñ5°oÑF×KÑKÓMÓNˆä�x‰xØ�<Ñ ×$Ñ$óóó
ð 	
r   c                 óR   — | j                   j                  || j                  ¬«       y)ai  
        Explain model predictions using provided data.

        References:
            https://dev.mysql.com/doc/heatwave/en/mys-hwaml-ml-explain-table.html
                A full list of supported options can be found under "ML_EXPLAIN_TABLE Options"

        Args:
            X: DataFrame for which predictions should be explained.

        Returns:
            DataFrame containing explanation details (feature attributions, etc.)

        Raises:
            DatabaseError:
                If provided options are invalid or unsupported,
                or if the model is not initialized, i.e., fit or import has not
                been called
                If a database connection issue occurs.
                If an operational error occurs during execution.

        Notes:
            Temporary input/output tables are cleaned up after explanation.
        r   N)r   Úexplain_predictionsr   )r   r   s     r   r2   z MyClassifier.explain_predictions§   s!   € ð6 	�‰×'Ñ'¨°4×3MÑ3MÐ'ÕNr   )NNNN)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r	   r   ÚstrÚdictr   r   ÚpdÚ	DataFramer-   Úndarrayr   r0   r2   r%   r   r   r   r   .   sæ   „ ñð6 %)Ø,0Ø04Ø04ñFà.ðFð ˜S‘MðFð $ D™>ð	Fð
  (¨™~ðFð  (¨™~óFðB/Ø�r—|‘| R§Z¡ZÐ/Ñ0ð/à	�‰ó/ð6"
Ø�r—|‘| R§Z¡ZÐ/Ñ0ð"
à	�‰ó"
ðHOØ�r—|‘| R§Z¡ZÐ/Ñ0ðOà	�‰ôOr   r   )r6   Útypingr   r   Únumpyr-   Úpandasr9   Úsklearn.baser   Úmysql.ai.ml.baser   Úmysql.ai.ml.modelr   Úmysql.ai.utilsr   Úmysql.connector.abstractsr	   r   r%   r   r   ú<module>rD      s6   ðñ:÷ #ã Û Ý (å *Ý %Ý $å =ôTO�= /õ TOr   