
    (HJj              	       V   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 d dlmZmZ ddlmZmZ ddlmZ dd	lmZmZ dd
lmZmZmZ defdZ	 dde
j<                  dedede fdZ!de
j<                  de"de
j<                  fdZ#de
j<                  de
j<                  fdZ$d Z%y)    N)Path)Dict)tree_flattentree_unflatten   )QuantizedSwitchLinearSwitchLinear)get_total_parameters   )DoRAEmbedding
DoRALinear)LoRAEmbedding
LoRALinearLoRASwitchLinearschedule_configc                 8   t        t        j                  | d         }| d   }|d   } || }| j                  dd      x}rY| j                  dd      }t        j                  j	                  |||      }t        j                  j                  ||g|dz   g      S |S )z?
    Build a learning rate schedule from the given config.
    name	argumentsr   warmupwarmup_initg        r   )getattropt
schedulersgetlinear_schedulejoin_schedules)r   schedule_fnr   
initial_lrbound_schedule_fnwarmup_stepsr   	warmup_fns           \/Users/ahmed/devFolder/claude-voice/.venv/lib/python3.12/site-packages/mlx_lm/tuner/utils.pybuild_scheduler#      s     #../&*ABK,I1J#Y/&**8Q77|7%))-=NN22\
	 ~~,,)*\A-=,>
 	
 !     model
num_layersconfiguse_dorac           	      &   fd}j                  dd      x1t               fd}| j                  D ]  }|j                  |        | j                  t	        |d       d D ]N  }|j                         D cg c]  \  }}|v s| ||      f }	}}|	s5|j                  t        |	             P | j                         D cg c]  \  }}|v s| ||      f }
}}|
r| j                  t        |
             yyc c}}w c c}}w )a  
    Convert some of the models linear layers to lora layers.

    Args:
        model (nn.Module): The neural network model.
        num_layers (int): The number of blocks to convert to lora layers
        starting from the last layer.
        config (dict): More configuration parameters for LoRA, including the
          rank, scale, and optional layer keys.
        use_dora (bool): If True, uses DoRA instead of LoRA.
          Default: ``False``
    c                 >   s)t        | d      r| j                  d   d   d         S t        | t        j                  t        j
                  f      rrt        nt        }nt        | t        t        f      r*r!t        t        |       j                   d      t        }n[t        | t        j                  t        j                  f      rrt         nt"        }n"t        dt        |       j                   d      |j%                  | d   d   d         S )	Nto_lorarankscaledropout)rr-   r.   z doesn't support DoRA yet.zCan't convert layer of type z to LoRA)hasattrr+   
isinstancennLinearQuantizedLinearr   r   r	   r   
ValueErrortype__name__r   	EmbeddingQuantizedEmbeddingr   r   	from_base)layer	LoRALayerr'   r(   s     r"   r+   z&linear_to_lora_layers.<locals>.to_lora9   s   GE95==.Woy) !   ebii););<=&.
JI.CDE DK$8$8#99S!TUU(Ib.C.CDE)1}I.tE{/C/C.DHM  ""Vn/9%	 # 
 	
r$   keysNc                     t         j                  t         j                  t        t        t         j
                  t         j                  f}t        |d      st        ||      rj                  |        y y )Nr+   )
r2   r3   r4   r	   r   r8   r9   r0   r1   add)pmtypesr=   s      r"   get_keys_for_loraz0linear_to_lora_layers.<locals>.get_keys_for_loraX   sU    		""%%%E q)$
1e(< )=r$   r   )r   setlayersapply_to_modulesmaxnamed_modulesupdate_modulesr   )r%   r&   r'   r(   r+   rC   lkrA   lora_layerslora_modulesr=   s     ``       @r"   linear_to_lora_layersrN   &   s   &
8 

64((1u
	 A01  \\3z1--/034??3DR3D41aT	71:3DR^K89 1
 160C0C0ES0E1dQ
O0ELS^L9:  S Ts   ?DDDDadapter_pathreturnc                    t        |      }|j                         st        d|       t        |dz  d      5 }t	        j
                  di t        j                  |      }ddd       t        dd      }|dk7  r&t        | |j                  |j                  |dk(  	       | j                  t        |d
z        d       | S # 1 sw Y   bxY w)a  
    Load any fine-tuned adapters / layers.

    Args:
        model (nn.Module): The neural network model.
        adapter_path (str): Path to the adapter configuration file.

    Returns:
        nn.Module: The updated model with LoRA layers applied.
    z!The adapter path does not exist: zadapter_config.jsonr/   Nfine_tune_typelorafulldora)r(   zadapters.safetensorsF)strict )r   existsFileNotFoundErroropenrB   SimpleNamespacejsonloadr   rN   r&   lora_parametersload_weightsstr)r%   rO   fidr'   rR   s        r"   load_adaptersrb   q   s     %L "CL> RSS	l22C	8C&&838 
9V%5v>N""$.		
 
s<*@@A%PL 
9	8s   )CCc                     g }| j                         D ]3  \  }}t        |t              s|j                  ||j                  f       5 t        |      dkD  r| j                  t        |             | S )z
    Remove the LoRA layers from the model.

    Args:
        model (nn.Module): The model with LoRA layers.

    Returns:
        nn.Module: The model without LoRA layers.
    r   )rH   r1   r   appendlinearlenrI   r   )r%   reset_layersr   modules       r"   remove_lora_layersri      sh     L++-ffj)v}} 56 . <1^L9:Lr$   c           	          t        |       dz  }t        d t        | j                               D              dz  }t	        d|dz  |z  dd|dd|dd       y )	Ng    .Ac              3   :   K   | ]  \  }}|j                     y w)N)size).0_vs      r"   	<genexpr>z-print_trainable_parameters.<locals>.<genexpr>   s     JItq!AFFIs   zTrainable parameters: d   z.3fz% (zM/zM))r
   sumr   trainable_parametersprint)r%   total_ptrainable_ps      r"   print_trainable_parametersrw      sr    "5)C/GJ|E,F,F,HIJJSP  

 +"3g"=s C DBwsm2	/r$   )F)&r\   rB   pathlibr   typingr   mlx.corecoremxmlx.nnr2   mlx.optimizers
optimizersr   	mlx.utilsr   r   models.switch_layersr   r	   utilsr
   rU   r   r   rS   r   r   r   r#   ModuleintboolrN   r`   rb   ri   rw   rW   r$   r"   <module>r      s           2 F ( + = =!D !0 	H;99H;H; H; 	H;V # ")) 8bii BII &r$   