
    (HJj/+                     |   d dl Z d dlZd dlZd dl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Zd dlZddlmZ ddlmZmZ ddlmZmZmZmZ ddlmZmZm Z m!Z! ddl"m#Z#m$Z$m%Z% ejL                  Z'e'jQ                  d	 ejR                  d
ejT                         e+d             i dddddddddi i i i i ddddd dddddddd d!d"d#d$d%d&d'dd(d)d*d+dd,d-dddd dd.d/d0d1dddd2Z,d3 Z-	 d;dej\                  d4efd5Z/dej\                  fd6Z0d;d4efd7Z1d8 Z2e3d9k(  r e4d:        e2        yy)<    N)Path   )get_reporting_callbacks)CacheDatasetload_dataset)TrainingArgsTrainingCallbackevaluatetrain)build_schedulelinear_to_lora_layersload_adaptersprint_trainable_parameters)_parse_sizeloadsave_configztag:yaml.org,2002:floatz^(?:
     [-+]?(?:[0-9][0-9_]*)\.[0-9_]*(?:[eE][-+]?[0-9]+)?
    |[-+]?(?:[0-9][0-9_]*)(?:[eE][-+]?[0-9]+)
    |\.[0-9_]+(?:[eE][-+][0-9]+)?
    |[-+]?[0-9][0-9_]*(?::[0-5]?[0-9])+\.[0-9_]*
    |[-+]?\.(?:inf|Inf|INF)
    |\.(?:nan|NaN|NAN))$z-+0123456789.modelzQwen/Qwen3-0.6br   Ffine_tune_typelora	optimizeradamoptimizer_configr   adamwmuonsgd	adafactordatazmlx-community/WikiSQLseed
num_layers   
batch_size   itersi  val_batches   learning_rategh㈵>steps_per_report
   steps_per_eval   resume_adapter_fileadapter_pathadapters
save_everyd   i  i      g        g      4@)rankdropoutscale)testtest_batchesmax_seq_lengthconfiggrad_checkpointgrad_accumulation_stepsclear_cache_thresholdlr_schedulelora_parametersmask_prompt	report_toproject_namec                     t        j                  d      } | j                  dt        d       | j                  dddd 	       | j                  d
t        d       | j                  dt        g dd       | j                  dt        g dd d       | j                  dddd 	       | j                  dt        d       | j                  dt        d       | j                  dt        d       | j                  dt        d       | j                  dt
        d       | j                  d t        d!       | j                  d"t        d#       | j                  d$t        d%       | j                  d&t        d'       | j                  d(t        d)       | j                  d*t        d+       | j                  d,dd-d 	       | j                  d.t        d/       | j                  d0t        d1       | j                  d2d3t        d4       | j                  d5dd6d 	       | j                  d7t        d8d9:       | j                  d;t        d d<:       | j                  d=t        d d>:       | j                  d?t        d@       | S )ANzLoRA or QLoRA finetuning.)descriptionz--modelz;The path to the local model directory or Hugging Face repo.)typehelpz--train
store_truezDo training)actionrD   defaultz--datazuDirectory with {train, valid, test}.jsonl files or the name of a Hugging Face dataset (e.g., 'mlx-community/wikisql')z--fine-tune-type)r   dorafullz4Type of fine-tuning to perform: lora, dora, or full.)rC   choicesrD   z--optimizerr   z>Optimizer to use for training: adam, adamw, sgd, or adafactor.)rC   rJ   rG   rD   z--mask-promptz)Mask the prompt in the loss when trainingz--num-layersz=Number of layers to fine-tune. Default is 16, use -1 for all.z--batch-sizezMinibatch size.z--iterszIterations to train for.z--val-batchesz@Number of validation batches, -1 uses the entire validation set.z--learning-ratezAdam learning rate.z--steps-per-reportz0Number of training steps between loss reporting.z--steps-per-evalz-Number of training steps between validations.z--grad-accumulation-stepsz;Number of steps to accumulate before each optimizer update.z--resume-adapter-filez?Load path to resume training from the given fine-tuned weights.z--adapter-pathz*Save/load path for the fine-tuned weights.z--save-everyz"Save the model every N iterations.z--testz'Evaluate on the test set after trainingz--test-batchesz8Number of test set batches, -1 uses the entire test set.z--max-seq-lengthzMaximum sequence length.z-cz--configz3A YAML configuration file with the training optionsz--grad-checkpointz0Use gradient checkpointing to reduce memory use.z--clear-cache-thresholdr   z>Clear the allocator cache between steps if it grows too large.)rC   rG   rD   z--report-tozDServices to report logs to ('wandb', 'swanlab', or 'wandb,swanlab').z--project-namezEProject name for logging. Defaults to the name of the root directory.z--seedzThe PRNG seed)argparseArgumentParseradd_argumentstrintfloatr   )parsers    U/Users/ahmed/devFolder/claude-voice/.venv/lib/python3.12/site-packages/mlx_lm/lora.pybuild_parserrS   Q   s(   $$1LMF
J   	   H	   (C	   =M   8	   L  
 S7HI
	2LM
O  
 )<QR
?  
 <  
 #J  
 N  
 9  
 1  
 6	   G  
 '  
 B	   ?	   !M	   S	   T	   sAM    training_callbackc                 ~   t         j                  j                  | j                         |j                          | j                  t        |j                        kD  r/t        d| j                   dt        |j                         d      | j                  dk(  rA|j                  t        | j                  d       d  D ]  }|j                           d | _        nW| j                  dv r1t        || j                  | j                  | j                  dk(         nt        d	| j                         | j                  5t        d
| j                          |j                  | j                  d       t!        |       t#        | j$                        }|j'                  dd       |dz  }t)        t+        |       |dz         t-        | j.                  | j0                  | j2                  | j4                  | j6                  | j8                  || j:                  | j<                  | j>                  
      }| j@                  rtC        | j@                        n| jD                  }	| jF                  jI                         }
| jJ                  jM                  |
i       }|
dk(  rtN        jP                  }nf|
dk(  rtN        jR                  }nP|
dk(  rtN        jT                  }n:|
dk(  rtN        jV                  }n$|
dk(  rtN        jX                  }nt        d|
        |dd|	i|}t[        |||t]        |      t]        |      |       y )NzRequested to train z layers but the model only has z layers.rI   r   )r   rH   rH   )use_doraz Received unknown fine-tune-type z Loading fine-tuned weights from F)strictT)parentsexist_okzadapters.safetensorszadapter_config.json)
r"   r$   r%   r(   r*   steps_per_saveadapter_filer7   r9   r:   r   r   r   r   r   zUnsupported optimizer: r'   )r   argsr   train_datasetval_datasetrU    )/mxrandomr   freezer    lenlayers
ValueErrorr   maxunfreezer=   r   r,   printload_weightsr   r   r-   mkdirr   varsr   r"   r$   r%   r(   r*   r/   r7   r9   r:   r<   r   r'   r   lowerr   getoptimAdamAdamWMuonSGD	Adafactorr   r   )r]   r   	train_set	valid_setrU   lr-   r\   training_argslroptimizer_namer   	opt_classopts                 rR   train_modelr}      s    IINN499	LLNU\\**!$//!2 3&&)%,,&7%8B
 	

 f$s4??A6689AJJL :  $			 0	0OO  ))V3		
 ;D<O<O;PQRR +01I1I0JKL433EBu%))*Ltd3"88LT
L+@@A !??jj$$..**!**,, $ < <M .2-=-=((	)4CUCUB^^))+N,,00DJJ		7	"KK		6	!JJ		5	 II		;	&OO	2>2BCDD

9"
9(8
9C 
"9- ++rT   c                     t        |t        |      | j                  | j                  | j                        }t        j                  |      }t        d|dd|dd       y )N)r   datasetr"   num_batchesr7   z
Test loss z.3fz, Test ppl .)r
   r   r"   r6   r7   mathexpri   )r]   r   test_set	test_losstest_ppls        rR   evaluate_modelr   1  s[    X&??%%**I xx	"H	Jyo[#a
@ArT   c                 p   t         j                  j                  | j                         t        | j                  | j
                  | j                  t        |             }t        d       t        | j                  ddi      \  }}t        d       t        | |      \  }}}| j                  r2| j                  s&| j                  dk7  rIt        || j                         n2| j                  rt        d       t        | ||||       nt!        d	      | j                  rt        d
       t#        | ||       y y )N)r@   log_dirr8   zLoading pretrained modeltrust_remote_codeT)tokenizer_configzLoading datasets Trainingz.Must provide at least one of --train or --testTesting)nprb   r   r   r?   r@   r-   rl   ri   r   r   r   r5   r   r   r}   rf   r   )r]   rU   r   	tokenizerru   rv   r   s          rR   runr   ?  s    IINN499/&&!!Dz	 

$%DJJ:Mt9TUE9	
%1$	%B"Iy(yy"%!2!23	jD%I7HIIJJyyitUH- rT   c                  "   dt         j                  d<   t               } | j                         }|j                  }t        |      }|rkt        d|       t        |d      5 }t        j                  |t              }d d d        |j                         D ]  \  }}|j                  |d       |||<    t        j                         D ]  \  }}|j                  |d       |||<    t        t        j                   di |       y # 1 sw Y   xY w)NtrueTOKENIZERS_PARALLELISMzLoading configuration filerr`   )osenvironrS   
parse_argsr8   rl   ri   openyamlr   yaml_loaderitemsrn   CONFIG_DEFAULTSr   typesSimpleNamespace)rQ   r]   r8   filekvs         rR   mainr   ^  s    +1BJJ'(^FD[[F:D*F3&#$YYt[1F  LLNDAqxx4 (Q #
  %%'188At$DG ( %%& s   DD__main__zwCalling `python -m mlx_lm.lora...` directly is deprecated. Use `mlx_lm.lora...` or `python -m mlx_lm lora ...` instead.)N)5rK   r   r   rer   warningspathlibr   mlx.corecorera   mlx.nnnnmlx.optimizers
optimizersro   numpyr   r   tuner.callbacksr   tuner.datasetsr   r   tuner.trainerr   r	   r
   r   tuner.utilsr   r   r   r   utilsr   r   r   
SafeLoaderr   add_implicit_resolvercompileXlistr   rS   Moduler}   r   r   r   __name__ri   r`   rT   rR   <module>r      s     	 	         4 6 J J  2 1oo  ! !BJJ	 		 	$$U$ f$ 	$
 $ #$ A$ "$ !$  T!$" 2#$$ T%$& '$( c)$* 4+$, J-$. #/$0   !cDAG$NDX +/V99V
 (VrB		 B.!1 .>', z		H 	F rT   