
    Nj                    R   d dl 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	m
Z
mZ d dlZd dlZd dlm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 e	rd dlmZ d dlmZ  ej@                  e!      Z"dZ#dZ$ddZ%ddZ&	 	 	 	 	 	 	 	 ddZddZ' G d dejP                        Z)ddZ*y)    )annotationsN)Counterwraps)TYPE_CHECKINGAnyCallable)train_test_split)MLPClassifierMLPRegressor)make_pipeline)	Tokenizer)nn)StaticModelPipeline)BaseFinetuneableStaticModelForClassification*   )z[PAD]z<pad>c                    | j                   | j                   d   S | j                         }t        D ]  }|j                  |      }||c S  t        j                  d       y)zRGet a probable pad token by using the padding module and falling back to guessing.pad_idz,No known pad token found, using 0 as defaultr   )padding	get_vocab_KNOWN_PAD_TOKENSgetloggerwarning)	tokenizervocabtokentoken_ids       f/Users/ahmed/devFolder/Ultron/claude-voice/.venv/lib/python3.12/site-packages/model2vec/train/utils.pyget_probable_pad_token_idr"      sb    $  **!E"99U#O #
 NNAB    c           	        ddl m} | j                         }t        j                  j                  d      }| j                  }|j                  ||j                        }t        | |      rA| j                  }t        | j                  f| j                  z        }| j                  rdnd}n8|j                  ||      }t        | j                  f| j                  z        }d}|j!                  ||       t#        | j$                  D 	cg c]  }	t        |	t&        j(                        s|	! c}	      D ]  \  }
}|j*                  j-                         j/                         j1                         j2                  |j4                  |
<   |j6                  j-                         j/                         j1                         |j8                  |
<    | j                  |_        ||_        t?        |      }tA        ||      S c c}	w )z)Convert the model to an sklearn pipeline.r   r   r   )hidden_layer_sizeslogisticsoftmaxidentity)!model2vec.train.classifierr   to_static_modelnprandomRandomStateout_dimrandndim
isinstanceclasses_r   
hidden_dimn_layers
multilabelr   fit	enumerateheadr   LinearweightdetachcpunumpyTcoefs_biasintercepts_
n_outputs_out_activation_r   r   )modelr   static_modelrandom_staten_itemsXymlp_head
activationmoduleindexlayerpipelines                r!   to_pipelinerP   ,   s   G((*L99((,LmmG7L$4$45A%56NN U5E5E4G%..4XY#(#3#3Z
w0E4D4D3F3WX
LLA!

"d
fjQWY[YbYbFc6
"deu!&!4!4!6!:!:!<!B!B!D!F!F&+jj&7&7&9&=&=&?&E&E&GU# f  --H)HX&H|X66 #es   G:(G:c                    d}t        |t              rZt        |d   t        t        f      rAt	        |      }t        |j                               dk  rt        j                  d       d}n|}t        | ||dd|      S )zSplit the data.

    For single-label classification, stratification is attempted (if possible).
    For multilabel classification, a random split is performed.
    Nr      zCSome classes have fewer than 2 samples. Stratification is disabled.r   T)	test_sizerF   shufflestratify)
r1   liststrintr   minvaluesr   infosklearn_split)rH   rI   rS   stratify_datalabel_countss        r!   r
   r
   K   sp     M!Tz!A$c
;qz|""$%)KK]^ MMAT\ijjr#   c                0     t               d fd       }|S )z'Suppresses annoying lightning warnings.c                     t        j                         5  t        j                  dd        | i |cd d d        S # 1 sw Y   y xY w)Nignore	lightning)rL   )warningscatch_warningsfilterwarnings)argskwargsfuncs     r!   wrapperz,suppress_lightning_warnings.<locals>.wrapperc   s8    $$&##H[A(( '&&s	   ?A)rf   r   rg   r   returnr	   r   )rh   ri   s   ` r!   suppress_lightning_warningsrk   `   s"     4[) )
 Nr#   c                      e Zd ZdZddZy)	TipFilterz7logging filter to suppress tip messages from lightning.c                &    d|j                         vS )z'Filter out tip messages from lightning.u   💡 Tip)
getMessage)selfrecords     r!   filterzTipFilter.filtero   s    !2!2!444r#   N)rq   zlogging.LogRecordrj   bool)__name__
__module____qualname____doc__rr    r#   r!   rm   rm   l   s
    A5r#   rm   c                :    t        j                  d| z  dz
         S )zInvert a sigmoid.   )torchlog)xs    r!   logitr~   t   s    IIq1uk"""r#   )r   r   rj   rX   )rD   z1'BaseFinetuneable | StaticModelForClassification'rj   r   )rH   z	list[str]rI   rV   rS   floatrj   z'tuple[list[str], list[str], list, list])rh   r	   rj   r	   )r}   torch.Tensorrj   r   )+
__future__r   loggingrc   collectionsr   	functoolsr   typingr   r   r	   r=   r+   r{   sklearn.model_selectionr
   r\   sklearn.neural_networkr   r   sklearn.pipeliner   
tokenizersr   r   model2vec.inferencer   model2vec.train.baser   r)   r   	getLoggerrt   r   _DEFAULT_RANDOM_SEEDr   r"   rP   rk   Filterrm   r~   rx   r#   r!   <module>r      s    "     / /   E > *    35G 
		8	$ & 7>kkk k -	k*	5 5#r#   