
    i!}                        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Zd dlZd dlm	Z	 d dl
mZ d dlmZmZmZmZmZmZmZmZ d dlmZ d dlmZ  ej4                  dd      j7                         dk(  r	 d dlmZ nd dlmZ  ej@                  ejB                  d
       d dl"m#Z#m$Z$m%Z%m&Z& ddl'm(Z( ddl'm)Z* dddddddddd	Z+dZ,dejZ                  dejZ                  fdZ.dee/ejZ                  f   dee/ef   deee/ejZ                  f   ee/ef   f   fdZ0de1fdZ2d  Z3d! Z4	 	 dTd"e/d#ee/   d$ee/   de	fd%Z5d& Z6d'e	de1fd(Z7d)d*de2fd'e	d+e8d,e8d-eee/ef      d.ee1geeejr                     ef   f   deejr                  e1f   fd/Z:d0ejr                  d1e/dejr                  fd2Z;dTd3Z<	 	 	 	 	 	 dUd"e/d4eee/ef      d-eee/ef      d1ee/   d+e8d5e8d#ee/   deeejr                  e(f   eejr                  e(ee/ef   f   f   fd6Z)	 	 	 dVd7eejz                  j|                     d8eejz                  j|                     d5e8fd9Z?dWd:Z@e,fde1d;eAdeBfd<ZCd=ee/e	f   d>ee/e	df   fd?ZDd=e/d@e/fdAZEd)dBdCee/e	f   d0ejr                  dDe8ddfdEZF	 	 dXd0ejr                  de1dFeeA   dGeeA   dHe/dIeee/ejr                  gee8e1f   f      deejr                  e1f   fdJZGd0ejr                  dejr                  fdKZHde1dLee/e	f   ddfdMZI	 dYdNee/e	f   dOee/e	f   d0ejr                  dPe(dee/ef   dDe8fdQZJdR ZKd0ejr                  de8fdSZLy# e$ r	  ed	      w xY w)Z    N)Path)dedent)AnyCallableDictListOptionalTupleTypeUnionMLXLM_USE_MODELSCOPEFalsetrue)snapshot_downloadz/Run `pip install modelscope` to use ModelScope.)i   i   )tree_flattentree_maptree_reducetree_unflatten   )TokenizerWrapper)loadllamamistral3phixtralmambadeepseek_v3qwen2_vlminimax)	mistralllavazphi-msftfalcon_mambajoyai_llm_flashkimi_k2
qwen2_5_vl
minimax_m2iquestcoder   qweightreturnc                     d}d|z  }| j                   \  }}||z  }d|z  dz
  }t        j                  g d      |z  }| d   |z	  |z  }|j                  ||      S )N       r   )r   r+   r   r'               ).N)shapemxarrayreshape)	r(   bitspack_factorout_features	packed_inin_featuresmaskshiftsunpackeds	            a/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/mlx_lm/utils.py_unpack_awq_weightsr>   <   sp    D*K%mmL)k)KI?DXX./$6F	"f,4HL+66    weightsquantization_configc                 |   |j                  dd      }|dk7  rt        d|d      |j                  dd      }i }t        | j                               D ]  j	                  d      rt        d d	      j	                  d
      rd d }| | d
   }| d}| d}| |   }	d|z  }
|j
                  \  }}||
z  }||z  }t        |      }|j                  }||
z  }|j                  |||
      }t        j                  |
      |z  }|j                  t        j                        |z  j                  d      j                  t        j                        }t        j                  |	j                        }	|| v r@| |   }t        |      }|j                  }|j                  t        j                         |	z  }n<d|dz
  z  }t        j                   |	j
                  | t        j                        |	z  }||| d<   |	|| d<   |j                  |	j"                        || d<   |	j"                  }t%        fddD              r|    |<    |j'                         D ]H  \  }}t        j(                  |j"                  t        j*                        s5|j                        ||<   J ||d}||fS )Nr5   r+   z
Only bits=z& is supported for AutoAWQ/GPTQ models.
group_size   z.g_idxzFound z in weights. Models with non-contiguous group indices (g_idx) are not currently supported. Please use a model without g_idx or re-quantize the model using mlx_lm.convert..qweighti.scales.qzerosr,   )axisr   )dtypez.weightz.biasesc              3   @   K   | ]  }j                  |        y wN)endswith).0suffixkeys     r=   	<genexpr>z)_transform_awq_weights.<locals>.<genexpr>   s      
/QVCLL /Qs   )rE   rG   rF   )rC   r5   )get
ValueErrorlistkeysrM   r1   r>   Tr4   r2   arangeastypeuint32sum
contiguousfloat32fullrJ   anyitems
issubdtypefloating)r@   rA   r5   rC   new_weightsprefixr(   
scales_key
qzeros_keyscalesr6   r9   
packed_outr7   n_groupsunpacked_weightr8   repackedr;   weightqzerosunpacked_zerosbiases
zero_pointmodel_dtypekwmlx_quantizationrP   s                               @r=   _transform_awq_weightsrt   G   s    ""61-Dqy;'MNOO$((s;JKGLLN#<<! A A  <<
#"XF12G"87+J"87+JZ(F *K&-mm#K%3L"j0H 2':O-//O ${2I&..|YTHYY{+d2F+v5:::CJJ299U  ]]688,F W$ , "5V!<!/!1!1
 )//

;;fD 4!8_

{"**MPVV.4K6('*+.4K6('*+.4mmFLL.IK6('*+ ,,K 
/Q
 
  's|KG $J !!#1=="++.XXk2KN $
 !
 (((r?   configc                     | d   }t         j                  ||      }	 t        j                  d|       }|j                  |j                  fS # t        $ r d| d}t        |      w xY w)z
    Retrieve the model and model args classes based on the configuration.

    Args:
        config (dict): The model configuration.

    Returns:
        A tuple containing the Model class and the ModelArgs class.
    
model_typezmlx_lm.models.zModel type z not supported.)MODEL_REMAPPINGrR   	importlibimport_moduleImportErrorrS   Model	ModelArgs)ru   rw   archmsgs       r=   _get_classesr      sz     %J $$Z<J&&
|'DE
 ::t~~%%	  J<7os   A A(c                 j    t        | j                         d       }d t        fd|D              S )Nc                 6    t        | t        j                        S rL   )
isinstancennModule)ms    r=   <lambda>z&get_total_parameters.<locals>.<lambda>   s    
1bii0Hr?   is_leafc                    t        | d      rMt        | d      sdn| j                  j                  }|| j                  j                  dz  | j                  z  z   S t        d t        | j                               D              S )Nr5   biasr   r,   c              3   :   K   | ]  \  }}|j                     y wrL   )size)rN   _vs      r=   rQ   z8get_total_parameters.<locals>.nparams.<locals>.<genexpr>   s     C&Bda166&Bs   )hasattrr   r   rk   r5   rZ   r   
parameters)r   ns     r=   nparamsz%get_total_parameters.<locals>.nparams   sa    1f F+Aqxx}}r)QVV333Cl1<<>&BCCCr?   c              3   4   K   | ]  \  }} |        y wrL    )rN   r   r   r   s      r=   rQ   z'get_total_parameters.<locals>.<genexpr>   s     3ldawqzls   )r   leaf_modulesrZ   )modelr   r   s     @r=   get_total_parametersr      s5    &HLD 3l333r?   c                 D    t        d | d      }t        |       }|dz  |z  S )Nc                 X    t        |t        j                        r| |j                  z   S | S rL   )r   r2   r3   nbytes)accxs     r=   r   z)compute_bits_per_weight.<locals>.<lambda>   s!    Arxx)@sQXX~IcIr?   r      )r   r   )r   model_bytesmodel_paramss      r=   compute_bits_per_weightr      s/    I5RSK (.L?\))r?   path_or_hf_reporevisionallow_patternsc                 z    t        |       }|j                         s|xs g d}t        t        | ||            }|S )a  
    Ensures the model is available locally. If the path does not exist locally,
    it is downloaded from the Hugging Face Hub.

    Args:
        path_or_hf_repo (str): The local path or Hugging Face repository ID of the model.
        revision (str, optional): A revision id which can be a branch name, a tag, or a commit hash.

    Returns:
        Path: The local file path.
    )	*.jsonmodel*.safetensors*.pytokenizer.model
*.tiktokentiktoken.model*.txt*.jsonl*.jinja)r   r   )r   existsr   )r   r   r   
model_paths       r=   	_downloadr      sN      o&J' 

 
,
 !-

 r?   c                 .    t        t        | d            S )NT)local_files_only)r   r   )hf_repos    r=   hf_repo_to_pathr      s    !'DABBr?   r   c                 x   t        | dz  d      5 }t        j                  |      }d d d        | dz  }|j                         rFi }	 t        |d      5 }t        j                  |      }d d d        |j                  dd      x}r|d<   S # 1 sw Y   fxY w# 1 sw Y   0xY w# t        j                  $ r Y Hw xY w)Nconfig.jsonrgeneration_config.jsoneos_token_idF)openjsonr   r   JSONDecodeErrorrR   )r   fru   generation_config_filegeneration_configr   s         r=   load_configr      s    	j=(#	.!1 
/ (*BB$$&	,c2a$(IIaL! 3
 -00GG<G%1F>"M 
/	. 32## 		s5   BB# B(B# BB B# #B98B9FTlazystrictmodel_configget_model_classesc                    t        |       |j                  |       t        j                  t        | dz              }|s|rt	        d|        i |D ]&  }j                  t        j                  |             ( j                  d      x}vt        j                  j                  d| |z        }t        j                  j                  |      }	|j                  j                  |	       |	j                  |	j                  }}
n |      \  }
}dvrj                  di       }d|v r|d   d<   |j!                        } |
|      t#        d	      rj%                        fd
}j                  dd      x}	 ||       nj                  dd      x}r{|d   }|dk(  rddlm}  ||      na|dk(  rdddd}|d<   |d<    ||       nC|dk(  rdddd}|d<   |d<    ||       n%|dv r!t+        |      \  }|d<   |d<    ||       j                  dd      rHd }t-        |j/                         t0        j2                  j4                        }j7                  |       j9                          j;                  t=        j?                               |       |s#t        j8                  jA                                fS )aB  
    Load and initialize the model from a given path.

    Args:
        model_path (Path): The path to load the model from.
        lazy (bool): If False eval the model parameters to make sure they are
            loaded in memory before returning, otherwise they will be loaded
            when needed. Default: ``False``
        strict (bool): Whether or not to raise an exception if weights don't
            match. Default: ``True``
        model_config (dict, optional): Optional configuration parameters for the
            model. Defaults to an empty dictionary.
        get_model_classes (Callable[[dict], Tuple[Type[nn.Module], Type]], optional):
            A function that returns the model class and model args class given a config.
            Defaults to the ``_get_classes`` function.

    Returns:
        Tuple[nn.Module, dict[str, Any]]: The loaded and initialized model and config.

    Raises:
        FileNotFoundError: If the weight files (.safetensors) are not found.
        ValueError: If the model class or args class are not found or cannot be instantiated.
    Nr   zNo safetensors found in 
model_filecustom_model)ru   rA   text_configsanitizec           	      r    fd}t        j                  | d   | d   | j                  dd      |       y )Nc                 J    | d   v rd   |    S t        |d      sy|  dv S )Nquantizationto_quantizedFrF   )r   )pr   ru   r@   s     r=   class_predicatez6load_model.<locals>._quantize.<locals>.class_predicateQ  s>    F>**n-a001n-S=G++r?   rC   r5   modeaffine)rC   r5   r   r   )r   quantizerR   )r   r   ru   r   r@   s     r=   	_quantizezload_model.<locals>._quantizeP  s<    	, 	#L1f%!!&(3+	
r?   r   Fquant_methodbitnetr   )bitnet_quantizemxfp4r,   r+   rC   r5   r   zcompressed-tensorsr   )awqgptqquantize_activationsc                 j   t        | t        j                        r| j                  dvrt	        d      | j                  dd      rt	        d      | j                  j                  \  }}|d| j                  z  z  }t        j                  ||| j                  | j                  | j                        S | S )N)nvfp4mxfp8z8Mode ({m.mode}) does not support activation quantizationr   Fz?Linear layer with bias does not support activation quantizationr,   )r   r   QuantizedLinearr   rS   rR   rk   r1   r5   QQLinearrC   )r   out_dimsin_dimss      r=   	_maybe_qqzload_model.<locals>._maybe_qq~  s    !R//066!33$R  55'$Y  %&HHNN!'2<'{{7HallAFFAFFSSr?   r   )r   )!r   updateglobstrFileNotFoundErrorr2   r   rR   ry   utilspec_from_file_locationmodule_from_specloaderexec_moduler|   r}   	from_dictr   r   models.bitlinear_layersr   rt   r   r   r   r   	is_moduleupdate_modulesevalload_weightsrT   r_   r   )r   r   r   r   r   weight_fileswfr   specr~   model_classmodel_args_classr   
model_argsr   r   rA   r   r   r   leavesru   r   r@   s                        @@@r=   
load_modelr     s   < $Fl#99S.B!BCDLF"::, GHHGrwwr{#  jj..
;~~55#
 ~~..t4%(,

DNN%(9(H%%F*jj3 K/,78M,NF()!++F3J
#Euj!..)
" 

>488E, &

+@% H	H		H*>:8#@#E+>?EW$*,aIL%1F>",8F()l#11*,aJL%1F>",8F()l#_,$:7DW$X!G\%1F>",8F()l#zz(%0	  )U%7%7%9299CVCVWV$	JJL	tGMMO,V<
  "#&=r?   r   adapter_pathc                      ddl m}  || |      S )Nr   )load_adapters)tuner.utilsr   )r   r   _load_adapterss      r=   r   r     s    <%..r?   c                 <    t        | g d      } t        | ||      S )z`Load a huggingface tokenizer and try to infer the type of streaming
    detokenizer to use.
    r   r   r   r   r   r   r   r   r   eos_token_ids)r   _load_tokenizer)r   tokenizer_config_extrar  s      r=   load_tokenizerr    s.     	
J # r?   tokenizer_configreturn_configc                     t        | |      }t        |||      \  }}	|t        ||      }|j                          t	        |||	j                  dd            }
|r||
|	fS ||
fS )aT  
    Load the model and tokenizer from a given path or a huggingface repository.

    Args:
        path_or_hf_repo (Path): The path or the huggingface repository to load the model from.
        tokenizer_config (dict, optional): Configuration parameters specifically for the tokenizer.
            Defaults to an empty dictionary.
        model_config(dict, optional): Configuration parameters specifically for the model.
            Defaults to an empty dictionary.
        adapter_path (str, optional): Path to the LoRA adapters. If provided, applies LoRA layers
            to the model. Default: ``None``.
        lazy (bool): If ``False`` eval the model parameters to make sure they are
            loaded in memory before returning, otherwise they will be loaded
            when needed. Default: ``False``
        return_config (bool: If ``True`` return the model config as the last item..
        revision (str, optional): A revision id which can be a branch name, a tag, or a commit hash.
    Returns:
        Union[Tuple[nn.Module, TokenizerWrapper], Tuple[nn.Module, TokenizerWrapper, Dict[str, Any]]]:
            A tuple containing the loaded model, tokenizer and, if requested, the model config.

    Raises:
        FileNotFoundError: If config file or safetensors are not found.
        ValueError: If model class or args class are not found.
    )r   )r   Nr   r  )r   r   r   r   r  rR   )r   r	  r   r   r   r
  r   r   r   ru   	tokenizers              r=   r   r     sx    H ?X>Jz4lKME6e\2

$FJJ~t4TI i''ir?   pipeline_grouptensor_groupc                     t        | g d      }t        |dd      \  }}t        |d      xr t        |j                  d      }t        |d      }||st	        d	      ||st	        d
      |s|st	        d      ||cxu rDn nA|rt
        j                  j                         }n |rt
        j                  j                         }||j                  j                  |       t        |dz  d      5 }	t        j                  |	      d   }
d d d        t               }t        |j                               D ]:  \  }}
j                  |d       d u x}rt	        d      |j!                  |
|          < t        | |       nt        |        t#        |ddi|j                  dd             }t        |dd      \  }}||j%                  |       ||j                  j                  |       t        j&                  |j                                t        j&                  t
        j                  j)                  t        j*                  d      t
        j,                               |r|||fS ||fS # 1 sw Y   gxY w)Nr  r  TF)r   r   r   pipelineshardzGThe model does not support pipelining but a pipeline_group was providedzMThe model does not support tensor parallelism but a tensor_group was providedz'The model does not support any shardingmodel.safetensors.index.jsonr   
weight_mapz<Pipeline loading is only supported for MLX converted models.trust_remote_coder   r  g      ?)stream)r   r   r   r   rS   r2   distributedinitr  r   r   r   setr   r   rR   addr  r  r   all_sumr3   cpu)repor  r  r
  r   r   ru   has_pipelininghas_tensor_parallelfidweight_indexlocal_filesrq   r   	file_namer  s                   r=   sharded_loadr#    sU    	
J  zUCME6UG,Qj1QN!%1!.U
 	
 (;[
 	
 "5BCC->>..0L^^002N !^, *==sCs99S>,7L D e !1!1!34DAq(,,Q5==y= R  OOLO, 5 	${3$ 	d#jj6I
 *4>HE1L!!^,GGE GGBNN""288C="@Ai''iE DCs   5I33I=c                 V    t        | t        j                  j                         d |      S rL   )r#  r2   r  r  )r  r
  s     r=   pipeline_loadr%  D  s     bnn113T=IIr?   max_file_size_gbc                     |dz  }g }i d}}| j                         D ]@  \  }}||j                  z   |kD  r|j                  |       i d}}|||<   ||j                  z  }B |j                  |       |S )z
    Splits the weights into smaller shards.

    Args:
        weights (dict): Model weights.
        max_file_size_gb (int): Maximum size of each shard in gigabytes.

    Returns:
        list: List of weight shards.
       r   )r_   r   append)r@   r&  max_file_size_bytesshardsr  
shard_sizerq   r   s           r=   make_shardsr-  H  s     +b0FA:E1 #66MM%  "A:Eaahh
   MM%Mr?   pathhf_pathc                    ddl m}m} ||j                   |d            }n|j	                  |      }d|j
                  _        d|j
                  _        |j
                  j                  dg|j
                  _        n8d|j
                  j                  vr |j
                  xj                  dgz  c_        |t        |      |j
                  _
        d|_        |j                  t        j                  j                  | d	             y)
z
    Uploads the model to Hugging Face hub.

    Args:
        path (Union[str, Path]): Local path to the model.
        hf_path (Union[str, Path, None]): Path to the original Hugging Face model.
    r   )	ModelCardModelCardDataNen)languagemlxztext-generation 	README.md)huggingface_hubr1  r2  from_templater   datalibrary_namepipeline_tagtagsr   
base_modeltextsaveosr.  join)r.  r/  r1  r2  cards        r=   create_model_cardrD  `  s     9&&}d'CD~~g&"DII.DIIyy~~			diinn	$		5'!"7|		DIIIbggll4-.r?   upload_repoc                    ddl m}m}m} ddlm} |j                          t        |       dz  }|j                  |      }|j                  j                  }|d| d| d	| d| d
| d}	nd}	t        d| d|	 d| d      |_        |j                  |        |       }
|
j                  |d       |
j                  | |d       t!        d| d       y)z
    Uploads the model to Hugging Face hub.

    Args:
        path (str): Local path to the model.
        upload_repo (str): Name of the HF repo to upload to.
    r   )HfApir1  loggingr   )__version__r7  Nz
        This model [z](https://huggingface.co/z,) was
        converted to MLX format from [z!)
        using mlx-lm version **z**.
        r6  z
        # z	
        z
        ## Use with mlx

        ```bash
        pip install mlx-lm
        ```

        ```python
        from mlx_lm import load, generate

        model, tokenizer = load("av  ")

        prompt = "hello"

        if tokenizer.chat_template is not None:
            messages = [{"role": "user", "content": prompt}]
            prompt = tokenizer.apply_chat_template(
                messages, add_generation_prompt=True, return_dict=False,
            )

        response = generate(model, tokenizer, prompt=prompt, verbose=True)
        ```
        T)repo_idexist_okr   )folder_pathrJ  	repo_typez0Upload successful, go to https://huggingface.co/z for details.)r8  rG  r1  rH  r6  rI  set_verbosity_infor   r   r:  r>  r   r?  r@  create_repoupload_large_folderprint)r.  rE  rG  r1  rH  rI  	card_pathrC  r/  
provenanceapis              r=   upload_to_hubrU  z  s    :9 T
[(I>>)$Dii""G M!:;- H''.i/H	 R  +} -	
 
- 		 
" #. /		DI6 	IIi
'COOK$O7  
 
<[M
WXr?   donate_model	save_pathrW  c                   t        | t              rt        |       } | j                  dd       t	        t        |j                                     }t        |      }t        |      }|dkD  rdnd}t        d |j                         D              }|t        |      di d}|r*|j                  t        d	 |j                                      |j                          ~t        t        |            D ]g  }	||	   }
d
||	<   |j!                  |	dz   |      }| |z  }t#        j$                  t        |      |
ddi       |
j'                         D ]
  }||d   |<    ~
i t)        |d         D ci c]  }||d   |    c}|d<   t+        | dz  d      5 }t-        j.                  ||d       d
d
d
       y
c c}w # 1 sw Y   y
xY w)z?Save model weights and metadata index into specified directory.T)parentsrK  r   z"model-{:05d}-of-{:05d}.safetensorszmodel.safetensorsc              3   4   K   | ]  }|j                     y wrL   )r   )rN   r   s     r=   rQ   zsave_model.<locals>.<genexpr>  s     8'7!QXX'7s   )
total_sizetotal_parameters)metadatar  c                 ,    t        j                  g       S rL   )r2   r3   )r   s    r=   r   zsave_model.<locals>.<lambda>  s    r?   Nformatr5  )r^  r  r  rr   r+   indent)r   r   r   mkdirdictr   r   r-  lenrZ   valuesr   r   r   clearranger`  r2   save_safetensorsrU   sortedr   r   dump)rX  r   rW  r@   r+  shards_countshard_file_formatr\  
index_datair  
shard_name
shard_pathweight_namerq   r   s                   r=   
save_modelrs    s    )S!O	OOD4O0< 0 0 234G!Fv;L ! 	-   8w~~'788J % 4U ;
 J X4e6F6F6HIJ MMO3v;q	q	&--a!e\B
+

C
OUh=NO ::<K4>J|$[1 (   17z,7O0P 0P1:l#A&&0P J| 
i88#	>!			
 
?	>	  
?	>s   ,F3F88GrC   r5   r   quant_predicatec                 4  	
 d }t        j                  |      xs t        | dd       |||      \  }||d
dv rd	nd	
d<   	
fd}t        j                  | |||	       d   d
<   t        |       }t        d|dd       | fS )a  
    Applies quantization to the model weights.

    Args:
        model (nn.Module): The model to be quantized.
        config (dict): Model configuration.
        group_size (Optional[int]): Group size for quantization.
        bits (Optional[int]): Bits per weight for quantization.
        mode (str): The quantization mode.
        quant_predicate (Callable): A callable that decides how to quantize
          each layer based on the path. Accepts the layer `path` and the
          `module`. Returns either a bool to signify quantize/no quantize or
          a dict of quantization parameters to pass to `to_quantized`.

    Returns:
        Tuple: Tuple containing quantized model and config.
    c                 8    ddddd}||    \  }}|xs ||xs |fS )N)@   r+   )r,   r+   )   r+   )r,   r   )r   r   r   r   r   )r   rC   r5   mode_defaultsdefault_group_sizedefault_bitss         r=   defaults_for_modez)quantize_model.<locals>.defaults_for_mode  s=    	
 ,9+>(L//1EEEr?   rt  Nr   r   TFc                     t        |d      sy|j                  j                  d   z  dk7  ryd}	 | |      }t        |t              r
|d   | <   |S r
|rd   | <   |S )Nr   FrH   r   Tr   )r   rk   r1   r   rd  )r.  modulebool_or_paramsfine_grained_configrC   quant_paramsrt  quantized_configs      r=   wrapped_predicatez)quantize_model.<locals>.wrapped_predicate)  s    v~.==r"Z/14&,T6:Nnd+5C^,T2  !^5A^,T2r?   )r   r   rA   z[INFO] Quantized model with z.3fz bits per weight.)copydeepcopygetattrr   r   r   rQ  )r   ru   rC   r5   r   rt  r|  r  bpwr  r  r  s     `  `   @@@r=   quantize_modelr    s    4F }}V,%P8I4)PO(z4@J",dDIL)) ##+7(  KK) /?~.N*+
!%
(C	(S	1B
CD"""r?   c           	         ddl m}m} g }| j                         D ]  \  }}d|v }t	        |t
        j                        rt
        j                  }d|i}nAt	        |t
        j                        ri }t
        j                  }nt	        ||      rd|i}|}n{t        j                  |j                  |j                  |j                  |j                  |j                   |j"                        }	|	j$                  ddd   }
 ||
i |}|r|j&                  |_        |	|_        |j)                  ||f        t+        |      dkD  r| j-                  t/        |             | S )z
    Dequantize the quantized layers in the model.

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

    Returns:
        nn.Module: The model with dequantized layers.
    r   )QuantizedSwitchLinearSwitchLinearr   NrH   r   )models.switch_layersr  r  named_modulesr   r   r   LinearQuantizedEmbedding	Embeddingr2   
dequantizerk   rf   rn   rC   r5   r   r1   r   r)  re  r   r   )r   r  r  dequantize_layersnamer~  r   clskwargsrk   argsr   s               r=   dequantize_modelr  G  s8    J++-ffb001))Cd^F 5 56F,,C 56d^FCMMMMMMKKKK
 ||DbD!  [[AF  $+5 .8 !^,=>?Lr?   config_pathc                    | j                  dd       | j                  dd       d| v r| d   | d<   t        t        | j                                     } t	        |d      5 }t        j                  | |d       ddd       y# 1 sw Y   yxY w)	a  Save the model configuration to the ``config_path``.

    The final configuration will be sorted before saving for better readability.

    Args:
        config (dict): The model configuration.
        config_path (Union[str, Path]): Model configuration file path.
    _name_or_pathNvision_configr   rA   rr   r+   ra  )poprd  rj  r_   r   r   rk  )ru   r  r  s      r=   save_configr  u  sy     JJ%
JJ%(.~(>$% &()F 
k3	3		&#a( 
 		s   BB
dst_pathsrc_path_or_repor  c                 l   t        |      }|j                         s|}t        |      }nd }t        |       } t        | |d       t	        || dz         |j                  |        dD ]>  }t        j                  t        ||z              D ]  }	t        j                  |	|         @ t        | |       y )NTrV  r   )r  )r   r   )r   r   r   rs  r  save_pretrainedr   r   shutilr  rD  )
r  r  r   r  ru   rW  src_pathr   r   files
             r=   r@  r@    s     $%H??""7+H~HxT2H}$<=h'/IIc(Q,/0DKKh' 1 0 h(r?   c                     t        t        |       t        |            }t        |      D ]  }| |   ||   k7  s|c S  |S )a$  
    Calculates the length of the common prefix of two lists.

    Args:
        list1: The first list of strings.
        list2: The second list of strings.

    Returns:
        The length of the common prefix. Returns 0 if lists are empty
        or do not match at the first element.
    )minre  rh  )list1list2min_lenro  s       r=   common_prefix_lenr    sD     #e*c%j)G 7^8uQxH  Nr?   c                     	 t        j                  | j                        }d|j                  v S # t        t
        f$ r Y yw xY w)z
    Check if the model supports input_embeddings in its call signature.
    Args:
        model (nn.Module): The model to check.
    Returns:
        bool: True if the model supports input_embeddings, False otherwise.
    input_embeddingsF)inspect	signature__call__r   rS   	TypeError)r   r  s     r=   #does_model_support_input_embeddingsr    sC    %%enn5	!Y%9%999	" s   ,/ A A)NN)NNNFFN)NNF)F)r   N)T)Mr  r   ry   r  r   rA  resourcer  pathlibr   textwrapr   typingr   r   r   r   r	   r
   r   r   mlx.corecorer2   mlx.nnr   getenvlower
modelscoper   r{   r8  	setrlimitRLIMIT_NOFILE	mlx.utilsr   r   r   r   tokenizer_utilsr   r   r  rx   MAX_FILE_SIZE_GBr3   r>   r   rt   rd  r   r   r   r   r   r   boolr   r   r   r  r  Groupr#  r%  intrT   r-  rD  rU  rs  r  r  r  r@  r  r  r   r?   r=   <module>r     sm        	    	 	 	  299#W-335?M0 2   8))< 8 I I . 4 $
  7 7bhh 7Y)#rxx- Y)c3hY) 4RXXS#X./Y)x& &*4* # $&&sm& I& 
	&RCD T * -1HTJJ
J J 4S>*	J
  d299ot.C(D DEJ 299d?JZ/ /# /")) /4 26-1"&"1 1 tCH~.1  4S>*1  3-	1 
 1  1  sm1  	"))%
%&	"))%tCH~
5681 l 6:37	T R^^112T  2>>//0T  	T nJ 8H   D 0/E#t), /uS$_7M /4?Y ?Y# ?YL 	9
S$Y9
999
 	9

 
9
B OSL#99L#L# L# 3-	L#
 L# hRYY'7tTz9J'JKLL# 299d?L#^+BII +")) +\))sDy!) 
)@ )CI)CI&) 99)  	)
 cN) )84ryy T O  MKLLMs   2L? ?M