
    Nj-B                    R   d dl mZ d dlZd dlmZ d dlmZ d dlmZm	Z	 d dl
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m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m Z  d dl!m"Z" d dl#m$Z$m%Z%m&Z&m'Z'm(Z(  ejR                  e*      Z+ G d dejX                        Z- e	de-      Z.y)    )annotationsN)Sequence)TemporaryDirectory)AnyTypeVar)LightningModule)CallbackEarlyStopping)Encoding	Tokenizer)nn)pad_sequence)trange)StaticModelPipeline)PathLikeStaticModel)TextDataset)get_probable_pad_token_idlogitsuppress_lightning_warningsto_pipelinetrain_test_splitc            
          e Zd ZdZdZdddddddddd			 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 d! fd
Zd"dZd#dZd$dZe		 d%dd	 	 	 	 	 	 	 	 	 d&d       Z
e	dd	 	 	 	 	 	 	 	 	 d'd       Zd(dZ ej                         d)d       Zd*d+dZd,dZd-d.dZed/d       Zd0dZd1dZd2dZ	 	 	 	 	 	 	 	 	 	 	 	 d3dZe	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 d4d       Ze	 	 	 	 	 	 	 	 d5d       Zd-d6dZd7dZ	 	 	 	 	 	 	 	 	 	 	 	 d8d Z xZ S )9BaseFinetuneableval_lossmin   r      NFT)	
hidden_dimn_layersout_dimpad_idtoken_mappingweightsfreeze	normalizefreeze_weightsc                  t         |           || _        || _        |j                  d   | _        || _        || _        |
| _        || _	        || _
        | j                  j                  t        j                  k7  rMt        | j                  j                        }t        j!                  d| d       |j#                         | _
        |+t        j$                  |t        j&                        | _        n3t        j*                  t-        |      t        j&                        | _        t/        j0                  | j(                  d      | _        |	| _        t.        j4                  j7                  |j9                         | j2                  |      | _        | j=                         | _        || _         | jC                         | _"        || _#        y)	aw  Initialize a trainable StaticModel from a StaticModel.

        :param vectors: The embeddings of the staticmodel.
        :param tokenizer: The tokenizer.
        :param hidden_dim: The hidden dimension of the head.
        :param n_layers: The number of layers in the head.
        :param out_dim: The output dimension of the head.
        :param pad_id: The padding id. This is set to 0 in almost all model2vec models
        :param token_mapping: The token mapping. If None, the token mapping is set to the range of the number of vectors.
        :param weights: The weights of the model. If None, the weights are initialized to zeros.
        :param freeze: Whether to freeze the embeddings. This should be set to False in most cases.
        :param normalize: Whether to normalize the embeddings.
        :param freeze_weights: Whether to freeze the learned token weights.
           zYour vectors are zI precision, converting to to torch.float32 to avoid compatibility issues.N)dtypeFrequires_gradr%   padding_idx)$super__init__r"   r!   shape	embed_dimr   r    r&   r'   vectorsr*   torchfloat32strloggerwarningfloattensorint64r#   arangelenr   	Parameterr%   	Embeddingfrom_pretrainedclone
embeddingsconstruct_headhead_weightsconstruct_weightsw	tokenizer)selfr3   rH   r   r    r!   r"   r#   r$   r%   r&   r'   r*   	__class__s                e/Users/ahmed/devFolder/Ultron/claude-voice/.venv/lib/python3.12/site-packages/model2vec/train/base.pyr0   zBaseFinetuneable.__init__$   sR   : 	 q)$ ",<<.**+ENN#E7*st #==?DL$!&m5;;!OD!&c'l%++!ND\\$*<*<ER,,66w}}t{{hn6o'')	'')"    c                   | j                   t        | j                         }nEt        j                  t	        | j
                              j                         }d|| j                  <   t        j                  || j                         S )z$Construct the weights for the model.ir+   )rE   r   r4   zerosr=   r#   r9   r"   r   r>   r'   )rI   rG   s     rK   rF   z"BaseFinetuneable.construct_weights^   sc    ==$dmm$AC 2 234::<A$AdkkN||A1D1D-DEErL   c                D   g }| j                   dk(  r:|j                  t        j                  | j                  | j
                               nt        j                  | j                  | j                        t        j                         g}t        | j                   dz
        D ]O  }|j                  t        j                  | j                  | j                        t        j                         g       Q |j                  t        j                  | j                  | j
                               |D cg c]  }t        |t        j                        s|! }}|r|^ }}|D ]V  }t        j                  j                  |j                  d       t        j                  j                  |j                         X t        j                  j!                  |j                         t        j                  j                  |j                         t        j"                  | S c c}w )z$Constructs a simple classifier head.r   r)   relu)nonlinearity)r    appendr   Linearr2   r!   r   ReLUrangeextend
isinstanceinitkaiming_uniform_weightzeros_biasxavier_uniform_
Sequential)rI   modules_modulelinear_modulesinitiallasts          rK   rC   zBaseFinetuneable.construct_headg   sf   #%==ANN299T^^T\\BC 		$..$//:	G 4==1,-		$//4?? KRWWYWX . NN299T__dllCD/6XwV*VRYY:W&wX+NWd!((V(Lv{{+ " GG##DKK0GGNN499%}}g&& Ys   1HHc                   | j                         | _        t        j                  j	                  | j
                  j                         | j                  | j                        | _	        | j                         | _        | j                          y)z'Initialize the classifier for training.r-   N)rC   rD   r   r?   r@   r3   rA   r%   r"   rB   rF   rG   trainrI   s    rK   _initializezBaseFinetuneable._initialize   sd    '')	,,66LL $++ 7 
 '')

rL   tokenc                   |j                  dd      x}rt        j                  d       |}t        j                  ||      } | j
                  dd|i|S )z1Load the model from a pretrained model2vec model.
model_nameNz<The 'model_name' argument is deprecated. Use 'path' instead.ri   model )popr7   r8   r   r@   from_static_model)clspathrj   kwargsrl   rm   s         rK   r@   z BaseFinetuneable.from_pretrained   sY      L$77:7NNYZD++D>$s$$;5;F;;rL   )	pad_tokenc          	        t        j                  |j                        |_        |j                  t	        j
                  |j                        nd}t	        j
                  |j                        }|j                  |j                  j                         }nd}||j                  j                         |   }nt        |j                        } | d|||j                  ||d|S )z#Load the model from a static model.N)r3   r"   rH   r#   r$   rn   )np
nan_to_num	embeddingr$   r4   
from_numpyr#   tolistrH   	get_vocabr   )rq   rm   rt   rs   r$   embeddings_convertedr#   r"   s           rK   rp   z"BaseFinetuneable.from_static_model   s     --85:]]5N%""5==1TX$//@*!//668M M __..0;F.u?F 
(oo'
 
 	
rL   c                   || j                   k7  j                         }|j                  d      dz   }| j                  |   }| j	                  |      }| j
                  |   }t        j                  |      }||z  }t        j                  |dddddf   |      j                  d      }||dddf   z  }| j                  rt        j                  j                  |      S |S )aN  A forward pass and mean pooling.

        This function is analogous to `StaticModel.encode`, but reimplemented to allow gradients
        to pass through.

        :param input_ids: A 2D tensor of input ids. All input ids are have to be within bounds.
        :return: The mean over the input ids, weighted by token weights.
        r)   gؗҜ<N)r"   r9   sumr#   rB   rG   r4   sigmoidbmmsqueezer&   r   
functional)rI   	input_idsrN   lengthinput_ids_embeddingsembeddedrG   s          rK   _encodezBaseFinetuneable._encode   s     dkk)0021%#11)<??#78FF9MM!I99Qq$z]H5==a@fQWo->>==**844rL   c                d    | j                  |      }| j                  | j                  |            S )N)tokenizerD   r   )rI   Xr   s      rK   _encode_single_batchz%BaseFinetuneable._encode_single_batch   s(    MM!$	yyi011rL   c                    g }t        dt        |      ||       D ]F  }| j                  ||||z          }|j                  |j	                         j                                H t        j                  |d      S )z#Encode a single batch of input ids.r   )disable)axis)r   r=   r   rR   cpunumpyrv   concatenate)rI   r   
batch_sizeshow_progress_barpredbatchlogitss          rK   encodezBaseFinetuneable.encode   so    As1vz?P;PQE..q9K/LMFKK

**,- R ~~d++rL   c                J    | j                  |      }| j                  |      |fS )z<Forward pass through the mean, and a classifier layer after.)r   rD   )rI   r   encodeds      rK   forwardzBaseFinetuneable.forward   s$    ,,y)yy!7**rL   c                    | j                   j                  |d      }|D cg c]2  }t        j                  |j                  d|       j                         4 }}t        |d| j                        S c c}w )a"  Tokenize a bunch of strings into a single padded 2D tensor.

        Note that this is not used during training.

        :param texts: The texts to tokenize.
        :param max_length: If this is None, the sequence lengths are truncated to 512.
        :return: A 2D padded tensor
        Fadd_special_tokensNT)batch_firstpadding_value)rH   encode_batch_fastr4   Tensoridslongr   r"   )rI   texts
max_lengthr   encodingencoded_idss         rK   r   zBaseFinetuneable.tokenize   sk     #'.."B"B5]b"B"cjq*rjq^f5<<[j8Q+R+W+W+Yjq*rKTUU +ss   7A3c                B    | j                   j                  j                  S )zGet the device of the model.)rB   rZ   devicerg   s    rK   r   zBaseFinetuneable.device   s     %%,,,rL   c                8   t        j                         5  | j                  j                  }|j	                         j                         }t        j                  | j                        j	                         j                         }ddd       t              t              k(  r0||dddf   z  }t        |d| j                  | j                  d      S t        ||| j                  | j                  | j                  j                               S # 1 sw Y   xY w)z$Convert the model to a static model.N)r3   r$   rH   r&   r#   )r4   no_gradrB   rZ   r   r   r   rG   r=   r   rH   r&   r#   )rI   embrG   s      rK   to_static_modelz BaseFinetuneable.to_static_model   s    ]]_//((C'')//#Cdff%))+113A  q6SX!T'
"C...."  nnnn,,224
 	
 _s   A0DDc                    t        |       S )z)Convert the model to an sklearn pipeline.)r   rg   s    rK   r   zBaseFinetuneable.to_pipeline  s    4  rL   c           	         |It        t        t        d|dz  dz        d            }t        |dz        }t        j	                  d|       |S )Nr)             z#Batch size automatically set to %d.)intr   maxr7   info)rI   r   train_lengthbase_numbers       rK   _determine_batch_sizez&BaseFinetuneable._determine_batch_size  sN    c#a,*;)B"CRHIK[2-.JKK=zJrL   c                    |d u|d uk7  rt        d      |3|1t        |d         t        |d         k7  rt        d      |}|}|}|}	nt        |||      \  }}}}	||||	fS )Nz;Both X_val and y_val must be provided together, or neither.r   z4X_val and y_val must be of the same type as X and y.)	test_size)
ValueErrortyper   )
rI   r   yX_valy_valr   train_textstrain_labelsvalidation_textsvalidation_labelss
             rK   _check_val_splitz!BaseFinetuneable._check_val_split  s     5#45Z[[!2E!H~ad+ !WXXKL$ %M]^_abnwMxJK)<9J,l<MMMrL   c
           
        g }
|4t        | j                  | j                  |d      }|
j                  |       | j	                  |	t        |      |      \  }}t               5 }t        j                  |||
||||      }|j                  ||j                  d|      |j                  d|             |j                  j                  }t        j                  |d      }d d d        i }d	   j                         D ]  \  }}d
|v r|||j!                  d      <     | j#                  |       | j%                          y # 1 sw Y   axY w)NgMbP?)monitormodepatience	min_delta)
min_epochs
max_epochs	callbacksval_check_intervalcheck_val_every_n_epochacceleratordefault_root_dirT)shuffler   F)train_dataloadersval_dataloaders)weights_only
state_dictloss_functionzmodel.)r
   
val_metricearly_stopping_directionrR   _determine_val_check_intervalr=   r   plTrainerfitto_dataloadercheckpoint_callbackbest_model_pathr4   loaditemsremoveprefixload_state_dicteval)rI   ra   train_datasetval_datasetr   early_stopping_patiencer   r   r   validation_stepsr   callbackr   check_val_every_epochtempdirtrainerr   best_model_weightsr   weight_namerZ   s                        rK   _trainzBaseFinetuneable._train7  s]    %'	".$220	H X&484V4Vc-0*5
11  !Wjj%%##5(="!(G KK"/"="=dWa"="b + 9 9%T^ 9 _  
 &99IIO!&O$!O# "& 
#5l#C#I#I#KK+-=CJ{//9:	 $L 	Z(		7 "!s   "B EEc                d    d }d}| #||z  }d}d}||kD  rt        |||z        }d }||fS | }d }||fS )Nr)         )r   )r   r   r   r   r   n_train_batchestarget_checks_per_epochmin_train_steps_between_vals           rK   r   z.BaseFinetuneable._determine_val_check_intervalo  s~     *.,-#*j8O&'#*-' !<<%(/#'>>&" )-%
 "#888 "2$(!!#888rL   c           	     <   |dz  }d}g }t        dt        |      dd      D ]c  }||||z    D cg c]  }|d| 	 }	}| j                  j                  |	d      }
|j	                  |
D cg c]  }|j
                  d|  c}       e t        ||      S c c}w c c}w )	zPrepare a dataset.

        :param X: The texts.
        :param y: The labels.
        :param max_length: The maximum length of the input.
        :return: A TextDataset.
        
      r   zTokenizing data)descNFr   )r   r=   rH   r   rV   r   r   )rI   r   r   r   truncate_lengthr   	tokenized	batch_idxxr   r   r   s               rK   _prepare_datasetz!BaseFinetuneable._prepare_dataset  s     %r/
%'	3q646GHI23I	J@V2WX2WQQ'(2WEXnn66uQV6WGPHhll;J7PQ I
 9a((	 YPs   B)B
c                    |S )zTurn the labels into a tensor.rn   )rI   labelss     rK   _labels_to_tensorz"BaseFinetuneable._labels_to_tensor  s    rL   c                   | j                  |||||      \  }}}}	| j                  |      }
| j                  |	      }t        j                  d       | j	                  ||
      }t        j                  d       | j	                  ||      }||fS )NzPreparing train dataset.zPreparing validation dataset.)r   r  r7   r   r   )rI   r   r   r   r   r   r   r   r   r   y_tensory_val_tensorr   r   s                 rK   _create_datasetsz!BaseFinetuneable._create_datasets  s     JNI^I^q%	J
F%|5F )),7--.?@./--k8D34++,<lKk))rL   )r3   torch.TensorrH   r   r   r   r    r   r!   r   r"   r   r#   zlist[int] | Noner$   ztorch.Tensor | Noner%   boolr&   r  r'   r  returnNone)r  znn.Parameter)r  znn.Sequential)r  r	  )zminishlab/potion-base-32m)
rq   type[ModelType]rr   r   rj   
str | Noners   r   r  	ModelType)
rq   r
  rm   r   rt   r  rs   r   r  r  )r   r  r  r  )r   	list[str]r  r  )r   F)r   r  r   r   r   r  r  z
np.ndarray)r   r  r  z!tuple[torch.Tensor, torch.Tensor])i   )r   r  r   
int | Noner  r  )r  ztorch.device)r  r   )r  r   )r   r  r   r   r  r   )r   r  r   listr   list[str] | Noner   zlist | Noner   r9   r  z/tuple[list[str], list[str], Sequence, Sequence])ra   r   r   r   r   r   r   r   r   r  r   r  r   r  r   r6   r   r  r  r	  )r   r  r   r   r   r   r  ztuple[int | None, int | None])r   r  r   r  r   r   r  r   )r   r   r  r  )r   r  r   r   r   r  r   z
Any | Noner   r9   r  ztuple[TextDataset, TextDataset])!__name__
__module____qualname__r   r   r0   rF   rC   rh   classmethodr@   rp   r   r4   r   r   r   r   r   propertyr   r   r   r   r   r   r   staticmethodr   r   r  r  __classcell__)rJ   s   @rK   r   r       s   J$ *.'+$8# 8# 	8#
 8# 8# 8# 8# (8# %8# 8# 8# 8# 
8#tF'8  5< !	<<< 	<
 < 
< < 
 !%	

 
 	

 
 

 
86 U]]_2 2,+
V - -
2!NN N  	N
 N N 
9N2 !55 #5 !	5
 5 ",5 5 5 5 %5 
5 !5n 9$9479EH9	&9 9.)(** *  	*
 * * 
)*rL   r   r  )bound)/
__future__r   loggingcollections.abcr   tempfiler   typingr   r   lightning.pytorchpytorchr   r   rv   r4   r   lightning.pytorch.callbacksr	   r
   
tokenizersr   r   r   torch.nn.utils.rnnr   tqdmr   model2vec.inferencer   model2vec.modelr   r   model2vec.train.datasetr   model2vec.train.utilsr   r   r   r   r   	getLoggerr  r7   Moduler   r  rn   rL   rK   <module>r*     sx    "  $ '     - ? *  +  3 1 /  
		8	$R*ryy R*j K'78	rL   