Ë
    Gêñi)  ã            
       ó~  — d Z ddlmZmZ ddlmZ ddlmZmZ  e«       r
ddl	Z	ddl
mZ  ej                  e«      Zdad„ Z G d„ d	ej"                  «      Z	 	 	 dd
ee   dz  defd„Zde	j.                  dedefd„Zde	j.                  de	j.                  de	j.                  dedef
d„Z G d„ de«      Z G d„ de«      Zy)a¹  
Metal affine quantization integration for transformers.

This module provides:
  - ``MetalLinear``: a drop-in replacement for ``nn.Linear`` that stores weights
    as affine-quantized uint32 packed tensors and uses the ``quantization-mlx``
    Metal kernels for the forward pass.
  - ``replace_with_metal_linear``: walks a model and swaps every eligible
    ``nn.Linear`` with ``MetalLinear``.
  - ``MetalQuantize`` / ``MetalDequantize``: weight conversion operations that
    participate in the new ``WeightConverter`` pipeline.

Weight layout (transposed, matching ``affine_qmm_t``):
  - ``weight``: ``[N, K_packed]`` (``uint32``) -- K is the packed dimension.
  - ``scales``:  ``[N, K // group_size]`` (``float16 / bfloat16``)
  - ``qbiases``: ``[N, K // group_size]`` (same dtype as scales)

The kernel call is ``affine_qmm_t(x, weight, scales, qbiases, group_size, bits)``
which computes ``y = x @ dequant(weight).T``, identical to ``nn.Linear``.
é   )ÚConversionOpsÚ_IdentityOp)Úshould_convert_module)Úis_torch_availableÚloggingé    Nc                  ó†   — t         €	 ddlm}   | d«      a t         S t         S # t        $ r}t	        d|› d�«      |‚d}~ww xY w)z>Lazily load the quantization-mlx kernel from Hugging Face Hub.Né   )Ú
get_kernelz0kernels-community/mlx-quantization-metal-kernelsz9Failed to load the quantization-mlx kernel from the Hub: zm. Make sure you have `kernels` installed (`pip install kernels`) and are running on an Apple Silicon machine.)Ú_metal_kernelÚhub_kernelsr   Ú	ExceptionÚImportError)r   Úes     ún/var/www/pod-logistic/pod-ai/venv/lib/python3.12/site-packages/transformers/integrations/metal_quantization.pyÚ_get_metal_kernelr   3   s`   € ô Ðð		Ý/á&Ð'YÓZˆMô ÐŒ=Ðøô ò 	ÜØKÈAÈ3ð O?ð ?óð ð	ûð	ús   ˆ" ¢	A «;»A c                   ó‚   — e Zd ZdZdej
                  ddfdedededed	ef
d
„Zdej                  dej                  fd„Z
y)ÚMetalLinearzê
    A quantized linear layer that stores weights in affine uint32 packed format
    and uses the ``quantization-mlx`` Metal kernels for the forward pass.

    Parameters match ``nn.Linear`` with additional quantization metadata.
    Fé   é€   Úin_featuresÚout_featuresÚbiasÚbitsÚ
group_sizec                 ó:  — t         j                  j                  | «       || _        || _        || _        || _        d|z  }||z  }||z  }	|t        j                  k(  rAt        j                  t        j                  ||t        j                  ¬«      d¬«      | _        n2t        j                  t        j                  |||¬«      d¬«      | _        |t        j                  k(  rt        j                  nd }
t        j                  t        j                  ||	|
¬«      d¬«      | _        t        j                  t        j                  ||	|
¬«      d¬«      | _        |r.t        j                  t        j                  |«      «      | _        y | j!                  dd «       y )Né    )ÚdtypeF)Úrequires_gradr   )ÚnnÚModuleÚ__init__r   r   r   r   ÚtorchÚuint32Ú	ParameterÚzerosÚweightÚfloat32ÚscalesÚqbiasesr   Úregister_parameter)Úselfr   r   r   r   r   r   Úelems_per_intÚk_packedÚn_groupsÚscales_dtypes              r   r"   zMetalLinear.__init__Q   s'  € ô 	�	‰	×Ñ˜4Ô à&ˆÔØ(ˆÔØˆŒ	Ø$ˆŒà˜d™
ˆØ -Ñ/ˆØ *Ñ,ˆà”E—L‘LÒ ÜŸ,™,¤u§{¡{°<ÀÔQV×Q]ÑQ]Ô'^ÐnsÔtˆD�KäŸ,™,¤u§{¡{°<ÀÐTYÔ'ZÐjoÔpˆDŒKà(-´·±Ò(=”u—}’}À4ˆÜ—l‘l¤5§;¡;¨|¸XÈ\Ô#ZÐjoÔpˆŒÜ—|‘|¤E§K¡K°¸hÈlÔ$[ÐkpÔqˆŒáÜŸ™¤U§[¡[°Ó%>Ó?ˆD�Ià×#Ñ# F¨DÕ1ó    ÚinputÚreturnc                 óü  — | j                   j                  t        j                  k7  r5t        j
                  j                  || j                   | j                  «      S t        «       }|j                  || j                   | j                  j                  |j                  «      | j                  j                  |j                  «      | j                  | j                  «      }| j                  �|| j                  z   }|S ©N)r'   r   r#   r$   r    Ú
functionalÚlinearr   r   Úaffine_qmm_tr)   Útor*   r   r   )r,   r2   ÚkernelÚoutputs       r   ÚforwardzMetalLinear.forwards   s°   € Ø�;‰;×Ñ¤§¡Ò,Ü—=‘=×'Ñ'¨¨t¯{©{¸D¿I¹IÓFÐFä"Ó$ˆà×$Ñ$ØØ�K‰KØ�K‰K�N‰N˜5Ÿ;™;Ó'Ø�L‰L�O‰O˜EŸK™KÓ(Ø�O‰OØ�I‰Ió
ˆð �9‰9Ð Ø˜dŸi™iÑ'ˆFØˆr1   N)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r#   r$   ÚintÚboolr"   ÚTensorr<   © r1   r   r   r   I   sj   „ ñð Ø�l‰lØØñ 2àð 2ð ð 2ð ð	 2ð ð 2ð ó 2ðD˜UŸ\™\ð ¨e¯l©lô r1   r   Úmodules_to_not_convertÚpre_quantizedc           
      óž  — |j                   r| S |j                  }|j                  }d}| j                  «       D ]z  \  }}t	        ||«      sŒt        |t        j                  «      sŒ.|ri nddi}	t        d|j                  |j                  |j                  du||dœ|	¤Ž}
| j                  ||
«       d}Œ| |st        j                  d«       | S )a`  
    Replace every eligible ``nn.Linear`` with ``MetalLinear``.

    Args:
        model: the ``PreTrainedModel`` (on the meta device at this point).
        modules_to_not_convert: module names to leave untouched.
        quantization_config: the ``MetalConfig`` instance.
        pre_quantized: ``True`` when loading from a quantized checkpoint.
    Fr   N)r   r   r   r   r   Tz�You are loading a model with Metal quantization but no nn.Linear modules were found. Please double check your model architecture.rD   )Ú
dequantizer   r   Únamed_modulesr   Ú
isinstancer    ÚLinearr   r   r   r   Úset_submoduleÚloggerÚwarning)ÚmodelrE   Úquantization_configrF   r   r   Úhas_been_replacedÚmodule_nameÚmoduleÚmodule_kwargsÚ
new_modules              r   Úreplace_with_metal_linearrV   ‡   sæ   € ð ×%Ò%Øˆà×#Ñ#€DØ$×/Ñ/€JàÐà$×2Ñ2Ó4ò %Ñˆ�VÜ$ [Ð2HÔIØä�fœbŸi™iÕ(Ù"/™B°g¸t°_ˆMÜ$ð Ø"×.Ñ.Ø#×0Ñ0Ø—[‘[¨Ð,ØØ%ñð  ñˆJð ×Ñ ¨ZÔ8Ø $Ñð!%ñ$ Ü�‰ð;ô	
ð
 €Lr1   r'   r   r   c                 ó
  — | j                   \  }}d|z  }d|z  dz
  }||z  }| j                  «       j                  |||«      }|j                  d¬«      j                  }	|j                  d¬«      j                  }
|
|	z
  |z  j                  d¬«      }|	}||j                  d«      z
  |j                  d«      z  }|j                  «       j                  d|«      j                  t        j                  «      j                  ||«      }||z  }t        j                  ||t        j                  | j                  ¬«      }t        |«      D ]  }||d	d	…|d	|…f   ||z  z  z  }Œ |j                  t        j                  «      ||fS )
aP  
    Quantize a 2-D float weight ``[N, K]`` into packed uint32 + scales + biases.

    Returns ``(w_packed, scales, biases)`` with:
      - ``w_packed``: ``[N, K // (32 // bits)]`` uint32
      - ``scales``:   ``[N, K // group_size]`` float32/float16/bfloat16
      - ``biases``:   ``[N, K // group_size]`` float32/float16/bfloat16
    r   r
   éÿÿÿÿ)Údimg:Œ0âŽyE>)Úminr   ©r   ÚdeviceN)ÚshapeÚfloatÚreshaperZ   ÚvaluesÚmaxÚclampÚ	unsqueezeÚroundr9   r#   Úint32r&   r\   Úranger$   )r'   r   r   ÚNÚKr-   Úmax_valr/   Ú	w_groupedÚw_minÚw_maxr)   ÚbiasesÚw_intr.   Úw_packedÚis                    r   Ú_affine_quantize_tensorrq   ¹   sm  € ð �<‰<�D€A€qØ˜$‘J€MØ�D‰y˜A‰o€GØ�J‰€Hà—‘“×&Ñ& q¨(°JÓ?€IØ�M‰M˜bˆMÓ!×(Ñ(€EØ�M‰M˜bˆMÓ!×(Ñ(€Eà�u‰} Ñ'×.Ñ.°4Ð.Ó8€FØ€Fà˜×)Ñ)¨"Ó-Ñ-°×1AÑ1AÀ"Ó1EÑE€EØ�K‰K‹M×Ñ  7Ó+×.Ñ.¬u¯{©{Ó;×CÑCÀAÀqÓI€Eð �MÑ!€HÜ�{‰{˜1˜h¬e¯k©kÀ&Ç-Á-ÔP€HÜ�=Ó!ò =ˆØ�Eš!˜QÐ- Ð-Ð-Ñ.°4¸!±8Ñ<Ñ<‰ð=ð �;‰;”u—|‘|Ó$ f¨fÐ4Ð4r1   ro   r)   rm   c                 ó2  — | j                   d   }d|z  }d|z  dz
  }| j                   d   |z  }| j                  t        j                  «      }	t        j                  ||t        j
                  | j                  ¬«      }
t        |«      D ]%  }|	||z  z	  |z  j                  «       |
dd…|d|…f<   Œ' |
j                  |d|«      }||j                  «       j                  d«      z  |j                  «       j                  d«      z   }|j                  ||«      S )zv
    Dequantize a packed uint32 weight ``[N, K_packed]`` back to float.

    Returns a ``[N, K]`` float32 tensor.
    r   r   r
   r[   NrX   )r]   r9   r#   re   r&   r(   r\   rf   r^   r_   rc   )ro   r)   rm   r   r   rg   r-   ri   rh   Ú
w_packed_iÚw_flatrp   rj   Úw_deqs                 r   Ú_affine_dequantize_tensorrv   Ú   s  € ð 	�‰�qÑ€AØ˜$‘J€MØ�D‰y˜A‰o€GØ�‰�qÑ˜MÑ)€Aà—‘œUŸ[™[Ó)€JÜ�[‰[˜˜A¤U§]¡]¸8¿?¹?ÔK€FÜ�=Ó!ò UˆØ(2°t¸a±xÑ(@ÀGÑ'K×&RÑ&RÓ&TˆŠq�!Ð"�]Ð"Ð"Ò#ðUð —‘˜q " jÓ1€IØ˜Ÿ™›×0Ñ0°Ó4Ñ4°v·|±|³~×7OÑ7OÐPRÓ7SÑS€EØ�=‰=˜˜AÓÐr1   c                   ó&   — e Zd ZdZd„ Zdedefd„Zy)ÚMetalQuantizezÃ
    Quantize a full-precision weight tensor into (weight, scales, qbiases).

    Used during quantize-on-the-fly.  The float ``weight`` is replaced in-place
    by the packed uint32 tensor.
    c                 ó   — || _         y r5   ©Úhf_quantizer©r,   r{   s     r   r"   zMetalQuantize.__init__ù   ó
   € Ø(ˆÕr1   Ú
input_dictr3   c                 óÚ  — t        t        |j                  «       «      «      \  }}t        |t        «      r|d   n|}| j
                  j                  j                  }| j
                  j                  j                  }t        |||«      \  }}}	d|v r|j                  dd«      d   nd}
|
r|
› d�nd}|
r|
› d�nd}|j                  }||||j                  |«      ||	j                  |«      iS )	Nr   ú.r
   Ú z.scalesr)   z.qbiasesr*   )ÚnextÚiterÚitemsrJ   Úlistr{   rP   r   r   rq   Úrsplitr   r9   )r,   r~   ÚkwargsÚ
target_keyÚvaluer   r   ro   r)   rm   ÚbaseÚ	scale_keyÚbias_keyÚ
orig_dtypes                 r   ÚconvertzMetalQuantize.convertü   sê   € Ü ¤ j×&6Ñ&6Ó&8Ó!9Ó:Ñˆ
�EÜ& u¬dÔ3��a’¸ˆà× Ñ ×4Ñ4×9Ñ9ˆØ×&Ñ&×:Ñ:×EÑEˆ
ä#:¸5À*ÈdÓ#SÑ ˆ�&˜&à/2°jÑ/@ˆz× Ñ   aÓ(¨Ò+ÀbˆÙ(,�t�f˜GÑ$°(ˆ	Ù(,�d�V˜8Ñ$°)ˆà—[‘[ˆ
à˜Ø�v—y‘y Ó,Ø�f—i‘i 
Ó+ð
ð 	
r1   N)r=   r>   r?   r@   r"   ÚdictrŽ   rD   r1   r   rx   rx   ñ   s   „ ñò)ð
 $ð 
°Tô 
r1   rx   c                   óD   — e Zd ZdZd„ Zd	dededz  defd„Zed
d„«       Z	y)ÚMetalDequantizezÊ
    Dequantize (weight, scales, qbiases) back to a full-precision tensor.

    Used when ``dequantize=True`` is set in the config to fall back to a normal
    ``nn.Linear`` on devices without MPS.
    c                 ó   — || _         y r5   rz   r|   s     r   r"   zMetalDequantize.__init__  r}   r1   Nr~   Úfull_layer_namer3   c                 ó4  — | j                   j                  j                  }| j                   j                  j                  }t	        |«      dk  r||d   iS |d   d   }|d   d   }|d   d   }t        |||||«      }	||	j                  |j                  «      iS )Nr   zweight$r   r)   r*   )r{   rP   r   r   Úlenrv   r9   r   )
r,   r~   r“   r‡   r   r   Ú	quantizedr)   r*   ru   s
             r   rŽ   zMetalDequantize.convert  s¤   € Ø× Ñ ×4Ñ4×9Ñ9ˆØ×&Ñ&×:Ñ:×EÑEˆ
äˆz‹?˜QÒØ# Z°	Ñ%:Ð;Ð;à˜yÑ)¨!Ñ,ˆ	Ø˜HÑ% aÑ(ˆØ˜YÑ'¨Ñ*ˆä)¨)°V¸WÀjÐRVÓWˆØ §¡¨&¯,©,Ó!7Ð8Ð8r1   c                 ó   — t        «       S r5   )r   )r,   s    r   Ú
reverse_opzMetalDequantize.reverse_op*  s
   € ä‹}Ðr1   r5   )r3   r   )
r=   r>   r?   r@   r"   r�   ÚstrrŽ   Úpropertyr˜   rD   r1   r   r‘   r‘     s?   „ ñò)ñ9 $ð 9¸¸t¹ð 9ÐY]ó 9ð òó ñr1   r‘   )NNF)r@   Úcore_model_loadingr   r   Úquantizers.quantizers_utilsr   Úutilsr   r   r#   Útorch.nnr    Ú
get_loggerr=   rM   r   r   rK   r   r…   r™   rB   rV   rC   rA   rq   rv   rx   r‘   rD   r1   r   ú<module>r       sï   ðñ÷* <Ý ?ß /ñ ÔÛÝð 
ˆ×	Ñ	˜HÓ	%€à€òô,;�"—)‘)ô ;ð@ 04ØØñ	/à  ™I¨Ñ,ð/ð ó	/ðd5 E§L¡Lð 5¸cð 5Èó 5ðBØ�l‰lðØ$)§L¡LðØ:?¿,¹,ðØTWðØ_bóô.
�Mô 
ô@�mõ r1   