
    (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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 Z-dej\                  dej\                  fdZ/dee0ej\                  f   dee0ef   deee0ej\                  f   ee0ef   f   fdZ1de2fd Z3d! Z4d" Z5	 	 dVd#e0d$ee0   d%ee0   de	fd&Z6d' Z7d(e	de2fd)Z8d*d+de3fd(e	d,e9d-e9d.eee0ef      d/ee2geeejt                     ef   f   deejt                  e2f   fd0Z;d1ejt                  d2e0dejt                  fd3Z<dVd4Z=	 	 	 	 	 	 dWd#e0d5eee0ef      d.eee0ef      d2ee0   d,e9d6e9d$ee0   deeejt                  e(f   eejt                  e(ee0ef   f   f   fd7Z)	 	 	 dXdd8d9eej|                  j~                     d:eej|                  j~                     d6e9d5eee0ef      fd;Z@dYd<ZAe,fde2d=eBdeCfd>ZDd?ee0e	f   d@ee0e	df   fdAZEd?e0dBe0fdCZFd*dDdEee0e	f   d1ejt                  dFe9ddfdGZG	 	 dZd1ejt                  de2dHeeB   dIeeB   dJe0dKeee0ejt                  gee9e2f   f      deejt                  e2f   fdLZHd1ejt                  dejt                  fdMZIde2dNee0e	f   ddfdOZJ	 d[dPee0e	f   dQee0e	f   d1ejt                  dRe(dee0ef   dFe9fdSZKdT ZLd1ejt                  de9fdUZMy# e$ r	  ed	      w xY w)\    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   c                     dddddd}d}| D ]  }|j                         s|dk(  s n|dz  }  t        | d |       }| |d  j                         j                         }t	        |||   z        S )Ng    .Ag    eAr   )MGMBGB r   .)isdigitfloatstripupperint)xsizessplitxidigitssizes         V/Users/ahmed/devFolder/claude-voice/.venv/lib/python3.12/site-packages/mlx_lm/utils.py_parse_sizer;   <   s    Cs#1=EE

c	
  1Ve9FefI$$&Dvd#$$    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	            r:   _unpack_awq_weightsrR   H   sp    D*K%mmL)k)KI?DXX./$6F	"f,4HL+66r<   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 )NrJ   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.qzerosrA   )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   )rX   rZ   rY   )rV   rJ   )get
ValueErrorlistkeysr`   rF   rR   TrI   rG   arangeastypeuint32sum
contiguousfloat32fullr]   anyitems
issubdtypefloating)rS   rT   rJ   rV   new_weightsprefixr=   
scales_key
qzeros_keyscalesrK   rN   
packed_outrL   n_groupsunpacked_weightrM   repackedrP   weightqzerosunpacked_zerosbiases
zero_pointmodel_dtypekwmlx_quantizationrc   s                               @r:   _transform_awq_weightsr   S   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_REMAPPINGre   	importlibimport_moduleImportErrorrf   Model	ModelArgs)r   r   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 r_   )
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 )NrJ   biasr   rA   c              3   :   K   | ]  \  }}|j                     y wr_   )r9   )ra   _vs      r:   rd   z8get_total_parameters.<locals>.nparams.<locals>.<genexpr>   s     C&Bda166&Bs   )hasattrr   r9   r~   rJ   rm   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r_    )ra   r   r   r   s      r:   rd   z'get_total_parameters.<locals>.<genexpr>   s     3ldawqzls   )r   leaf_modulesrm   )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 r_   )r   rG   rH   nbytes)accr4   s     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   JSONDecodeErrorre   )r   fr   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)r   rT   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_quantizedFrY   )r   )pr   r   rS   s     r:   class_predicatez6load_model.<locals>._quantize.<locals>.class_predicate]  s>    F>**n-a001n-S=G++r<   rV   rJ   modeaffine)rV   rJ   r   r   )r   quantizere   )r   r   r   r   rS   s     r:   	_quantizezload_model.<locals>._quantize\  s<    	, 	#L1f%!!&(3+	
r<   r   Fquant_methodbitnetr   )bitnet_quantizemxfp4rA   r@   rV   rJ   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 quantizationrA   )r   r   QuantizedLinearr   rf   re   r~   rF   rJ   QQLinearrV   )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FileNotFoundErrorrG   r   re   r   utilspec_from_file_locationmodule_from_specloaderexec_moduler   r   	from_dictr   r   models.bitlinear_layersr   r   r   r   r   r   	is_moduleupdate_modulesevalload_weightsrg   rr   r   )r   r   r   r   r   weight_fileswfr   specr   model_classmodel_args_classr   
model_argsr   r   rT   r   r   r   leavesr   r   rS   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  re   )r   r  r   r  r   r  r   r   r   r   	tokenizers              r:   r   r     sx    H ?X>Jz4lKME6e\2

$FJJ~t4TI i''ir<   )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#        ||xs 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   kx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   rf   rG   distributedinitr!  r   r   r   setr   r   re   addr  r"  r  all_sumrH   cpu)repor  r  r  r  r   r   r   has_pipelininghas_tensor_parallelfidweight_indexlocal_filesr   r   	file_namer  s                    r:   sharded_loadr4    s[    	
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$ 70$7jj6I
 *4>HE1L!!^,GGE GGBNN""288C="@Ai''iE DCs   5I77Jc                 V    t        | t        j                  j                         d |      S r_   )r4  rG   r'  r(  )r-  r  s     r:   pipeline_loadr6  R  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   )rr   r   append)rS   r7  max_file_size_bytesshardsr"  
shard_sizer   r   s           r:   make_shardsr>  V  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-generationr-   	README.md)huggingface_hubrB  rC  from_templater   datalibrary_namepipeline_tagtagsr   
base_modeltextsaveosr?  join)r?  r@  rB  rC  cards        r:   create_model_cardrT  n  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   )HfApirB  loggingr   )__version__rG  Nz
        This model [z](https://huggingface.co/z,) was
        converted to MLX format from [z!)
        using mlx-lm version **z**.
        r-   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_pathrZ  	repo_typez0Upload successful, go to https://huggingface.co/z for details.)rH  rW  rB  rX  r-   rY  set_verbosity_infor   r   rJ  rN  r   rO  rP  create_repoupload_large_folderprint)r?  rU  rW  rB  rX  rY  	card_pathrS  r@  
provenanceapis              r:   upload_to_hubre    s    :9 T
[(I>>)$Dii""G M!:;- H''.i/H	 R  +} -	
 
- 		 
" #. /		DI6 	IIi
'COOK$O7  
 
<[M
WXr<   donate_model	save_pathrg  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)parentsr[  r   z"model-{:05d}-of-{:05d}.safetensorszmodel.safetensorsc              3   4   K   | ]  }|j                     y wr_   )r   )ra   r   s     r:   rd   zsave_model.<locals>.<genexpr>  s     8'7!QXX'7s   )
total_sizetotal_parameters)metadatar$  c                 ,    t        j                  g       S r_   )rG   rH   )r   s    r:   r   zsave_model.<locals>.<lambda>  s    r<   NformatrF  )rn  r$  r#  r   r@   indent)r   r   r   mkdirdictr   r   r>  lenrm   valuesr   r   r   clearrangerp  rG   save_safetensorsrh   sortedr   r   dump)rh  r   rg  rS   r<  shards_countshard_file_formatrl  
index_datair"  
shard_name
shard_pathweight_namer   r   s                   r:   
save_modelr    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rV   rJ   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@   )rA   r@   )   r@   )rA   r   )r   r   r   r   r   )r   rV   rJ   mode_defaultsdefault_group_sizedefault_bitss         r:   defaults_for_modez)quantize_model.<locals>.defaults_for_mode   s=    	
 ,9+>(L//1EEEr<   r  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   Fr[   r   Tr   )r   r~   rF   r   rt  )r?  modulebool_or_paramsfine_grained_configrV   quant_paramsr  quantized_configs      r:   wrapped_predicatez)quantize_model.<locals>.wrapped_predicate7  s    v~.==r"Z/14&,T6:Nnd+5C^,T2  !^5A^,T2r<   )r   r   rT   z[INFO] Quantized model with z.3fz bits per weight.)copydeepcopygetattrr   r   r   ra  )r   r   rV   rJ   r   r  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   Nr[   r   )models.switch_layersr  r  named_modulesr   r   r   LinearQuantizedEmbedding	EmbeddingrG   
dequantizer~   ry   r   rV   rJ   r   rF   r   r:  ru  r  r   )r   r  r  dequantize_layersnamer  r   clskwargsr~   argsr   s               r:   dequantize_modelr  U  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   rT   r   r@   rq  )poprt  rz  rr   r   r   r{  )r   r  r0  s      r:   save_configr    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 )NTrf  r   )r  )r   r   )r   r   r   r  r  save_pretrainedr   r   shutilr  rT  )
r  r  r   r  r   rg  src_pathr   r   files
             r:   rP  rP    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.
    )minru  rx  )list1list2min_lenr  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   rf   	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)Nr  r   r   r  r   rQ  resourcer  pathlibr   textwrapr   typingr   r   r   r   r	   r
   r   r   mlx.corecorerG   mlx.nnr   getenvlower
modelscoper   r   rH  	setrlimitRLIMIT_NOFILE	mlx.utilsr   r   r   r   tokenizer_utilsr   r   r  r   MAX_FILE_SIZE_GBr;   rH   rR   r   r   rt  r   r   r   r   r   r   boolr   r  r  r  r'  Groupr4  r6  r3   rg   r>  rT  re  r  r  r  r  rP  r  r  r   r<   r:   <module>r     s        	    	 	 	  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	V  26V R^^112V  2>>//0V  	V  tCH~.V r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 k  MKLLMs   2M M