+
    QV-j)  ã                   ó  € R t ^RIHtHt ^RIHt ^RIHtHt ]! 4       '       d   ^ RI	t	^ RI
Ht ]P                  ! ]4      tRsR t ! R R]P"                  4      tRR R	 lltR
 R ltR R lt ! R R]4      t ! R R]4      tR# )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ÚloggingNc                 óŽ   € \         f    ^RIHp  V ! R4      s \         # \         #   \         d   p\	        RT R24      ThRp?ii ; i)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     Ú}/Volumes/fast/ai/experiments/ui-tars-smoke/.venv/lib/python3.14/site-packages/transformers/integrations/metal_quantization.pyÚ_get_metal_kernelr   3   s`   € ô Òð		Ý/á&Ð'YÓZˆMô ÐŒ=Ðøô ô 	ÜØKÈAÈ3ð O?ð ?óð ð	ûð	ús   Š$ ¤A¯?¿Ac                   óh   a € ] tR t^It o RtR]P                  ^^€3V 3R lR lltV 3R lR ltRt	V t
R# )	Ú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.
Fc          
      ó8   <€ V ^8„  d   QhRS[ RS[ RS[RS[ RS[ /# )é   Úin_featuresÚout_featuresÚbiasÚbitsÚ
group_size)ÚintÚbool)ÚformatÚ__classdict__s   "€r   Ú__annotate__ÚMetalLinear.__annotate__Q   s=   ø€ ÷  2ñ  2áð 2ñ ð 2ñ ð	 2ñ ð 2ñ ñ 2ó    c                ój  € \         P                  P                  V 4       Wn        W n        WPn        W`n        ^ V,          pW,          pW,          p	V\        P                  8X  dC   \         P                  ! \        P                  ! W(\        P                  R7      RR7      V n        M3\         P                  ! \        P                  ! W!VR7      RR7      V n        V\        P                  8X  d   \        P                  MRp
\         P                  ! \        P                  ! W)V
R7      RR7      V n        \         P                  ! \        P                  ! W)V
R7      RR7      V n        V'       d2   \         P                  ! \        P                  ! V4      4      V n        R# V P!                  RR4       R# )é    )ÚdtypeF)Úrequires_gradNr   )Ú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&   Ú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§;¢;¨|È\Ô#ZÐjoÔpˆŒÜ—|’|¤E§K¢K°ÈlÔ$[ÐkpÔqˆŒçÜŸš¤U§[¢[°Ó%>Ó?ˆDŽIà×#Ñ# F¨DÖ1r   c                óN   <€ V ^8„  d   QhRS[ P                  RS[ P                  /# )r   ÚinputÚreturn)r'   ÚTensor)r   r   s   "€r   r   r   s   s#   ø€ ÷ ñ ™UŸ\™\ð ©e¯l©lñ r   c                ó  € V P                   P                  \        P                  8w  d5   \        P
                  P                  WP                   V P                  4      # \        4       pVP                  VV P                   V P                  P                  VP                  4      V P                  P                  VP                  4      V P                  V P                  4      pV P                  e   W0P                  ,           pV# ©N)r+   r"   r'   r(   r$   Ú
functionalÚlinearr   r   Úaffine_qmm_tr-   Útor.   r   r   )r0   r7   ÚkernelÚoutputs   &&  r   ÚforwardÚMetalLinear.forwards   s¬   € Ø�;‰;×Ñ¤§¡Ô,Ü—=‘=×'Ñ'¨¯{©{¸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Ò ØŸi™iÕ'ˆFØˆr   )r   r   r   r   r   r.   r-   r+   N)Ú__name__Ú
__module__Ú__qualname__Ú__firstlineno__Ú__doc__r'   r(   r&   rB   Ú__static_attributes__Ú__classdictcell__©r   s   @r   r   r   I   s1   ø‡ € ñð Ø�l‰lØØ÷ 2ò  2÷Dö r   r   c                óT   € V ^8„  d   QhR\         \        ,          R,          R\        /# )r   Úmodules_to_not_convertNÚpre_quantized)ÚlistÚstrr   )r   s   "r   r   r   ‡   s'   € ÷ /ñ /ä ¤�I¨Õ,ð/ô ñ	/r   c                óä  € VP                   '       d   V # VP                  pVP                  pRpV P                  4        F�  w  rx\	        Wq4      '       g   K  \        V\        P                  4      '       g   K:  V'       d   / MRR/p	\        RRVP                  RVP                  RVP                  RJRVRV/V	B p
V P                  Wz4       R	pK’  	  V'       g   \        P                  R
4       V # )aD  
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"   Nr   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.© )Ú
dequantizer   r   Únamed_modulesr   Ú
isinstancer$   ÚLinearr   r   r   r   Úset_submoduleÚloggerÚwarning)ÚmodelrM   Úquantization_configrN   r   r   Úhas_been_replacedÚmodule_nameÚmoduleÚmodule_kwargsÚ
new_modules   &&&&       r   Úreplace_with_metal_linearra   ‡   sý   € ð ×%×%Ð%Øˆà×#Ñ#€DØ$×/Ñ/€JàÐà$×2Ñ2Ö4ÑˆÜ$ [×IÒIÙä�fœbŸi™i×(Ô(ß"/™B°g¸t°_ˆMÜ$ñ Ø"×.Ñ.ðà#×0Ñ0ðð —[‘[¨Ð,ðð ð	ð
 &ðð  ñˆJð ×Ñ Ô8Ø $Òñ!  5÷$ Ü�‰ð;ô	
ð
 €Lr   c                óP   € V ^8„  d   QhR\         P                  R\        R\        /# )r   r+   r   r   ©r'   r9   r   )r   s   "r   r   r   ¹   s%   € ÷ 5ñ 5¤E§L¡Lð 5¼cð 5Ìñ 5r   c                ó`  € V P                   w  r4^ V,          p^V,          ^,
          pWA,          pV P                  4       P                  W7V4      pVP                  RR7      P                  p	VP                  RR7      P                  p
W©,
          V,          P                  RR7      pT	pWŒP                  R4      ,
          VP                  R4      ,          pVP                  4       P                  ^ V4      P                  \        P                  4      P                  W44      pWE,          p\        P                  ! W>\        P                  V P                  R7      p\        V4       F%  pWýRVRV13,          VV,          ,          ,          pK'  	  VP                  \        P                  4      W¼3# )a8  
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
)Údimg:Œ0âŽyE>)Úmin©r"   ÚdeviceºNNNNéÿÿÿÿ)ÚshapeÚfloatÚreshaperf   ÚvaluesÚmaxÚclampÚ	unsqueezeÚroundr?   r'   Úint32r*   rh   Úranger(   )r+   r   r   ÚNÚKr1   Úmax_valr3   Ú	w_groupedÚw_minÚw_maxr-   ÚbiasesÚw_intr2   Úw_packedÚis   &&&              r   Ú_affine_quantize_tensorr   ¹   sT  € ð �<‰<�D€AØ˜$•J€MØ�D�y˜A�o€GØ�€Hà—‘“×&Ñ& q°JÓ?€IØ�M‰M˜bˆMÓ!×(Ñ(€EØ�M‰M˜bˆMÓ!×(Ñ(€Eà�} Õ'×.Ñ.°4Ð.Ó8€FØ€Fà×)Ñ)¨"Ó-Õ-°×1AÑ1AÀ"Ó1EÕE€EØ�K‰K‹M×Ñ  7Ó+×.Ñ.¬u¯{©{Ó;×CÑCÀAÓI€Eð Õ!€HÜ�{Š{˜1¬e¯k©kÀ&Ç-Á-ÔP€HÜ�=Ö!ˆØ˜!˜QÐ- Ð-Ð-Õ.°4¸!µ8Õ<Õ<Šñ "ð �;‰;”u—|‘|Ó$ fÐ4Ð4r   c          
      ó�   € V ^8„  d   QhR\         P                  R\         P                  R\         P                  R\        R\        /# )r   r}   r-   r{   r   r   rc   )r   s   "r   r   r   Ú   s@   € ÷ ñ Ü�l‰lðÜ$)§L¡LðÜ:?¿,¹,ðÜTWðÜ_bñr   c                ó„  € V P                   ^ ,          p^ V,          p^V,          ^,
          pV P                   ^,          V,          pV P                  \        P                  4      p	\        P                  ! WX\        P
                  V P                  R7      p
\        V4       F/  pW”V,          ,	          V,          P                  4       V
RVRV13&   K1  	  V
P                  VRV4      pWÁP                  4       P                  R4      ,          VP                  4       P                  R4      ,           pVP                  WX4      # )zj
Dequantize a packed uint32 weight ``[N, K_packed]`` back to float.

Returns a ``[N, K]`` float32 tensor.
rg   ri   Nrj   )rk   r?   r'   rs   r*   r,   rh   rt   rl   rm   rq   )r}   r-   r{   r   r   ru   r1   rw   rv   Ú
w_packed_iÚw_flatr~   rx   Úw_deqs   &&&&&         r   Ú_affine_dequantize_tensorr…   Ú   sõ   € ð 	�‰�qÕ€AØ˜$•J€MØ�D�y˜A�o€GØ�‰�qÕ˜MÕ)€Aà—‘œUŸ[™[Ó)€JÜ�[Š[˜¤U§]¡]¸8¿?¹?ÔK€FÜ�=Ö!ˆØ(2¸aµxÕ(@ÀGÕ'K×&RÑ&RÓ&Tˆˆq�!Ð"�]Ð"Ð"Ó#ñ "ð —‘˜q " jÓ1€IØŸ™›×0Ñ0°Ó4Õ4°v·|±|³~×7OÑ7OÐPRÓ7SÕS€EØ�=‰=˜ÓÐr   c                   ó<   a € ] tR t^ñt o RtR tV 3R lR ltRtV tR# )Ú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                ó   € Wn         R # r;   ©Úhf_quantizer©r0   rŠ   s   &&r   r&   ÚMetalQuantize.__init__ù   ó   € Ø(Ör   c                ó&   <€ V ^8„  d   QhRS[ RS[ /# )r   Ú
input_dictr8   )Údict)r   r   s   "€r   r   ÚMetalQuantize.__annotate__ü   s   ø€ ÷ 
ñ 
¡$ð 
±Tñ 
r   c                ó  € \        \        VP                  4       4      4      w  r4\        V\        4      '       d
   V^ ,          MTpV P
                  P                  P                  pV P
                  P                  P                  p\        WFV4      w  rxp	RV9   d   VP                  R^4      ^ ,          MRp
V
'       d   V
 R2MRpV
'       d   V
 R2MRpVP                  pW7W¸P                  V4      WÉP                  V4      /# )é    Ú.Ú z.scalesr-   z.qbiasesr.   )ÚnextÚiterÚitemsrU   rO   rŠ   r[   r   r   r   Úrsplitr"   r?   )r0   r�   ÚkwargsÚ
target_keyÚvaluer   r   r}   r-   r{   ÚbaseÚ	scale_keyÚbias_keyÚ
orig_dtypes   &&,           r   ÚconvertÚMetalQuantize.convertü   sá   € Ü ¤ j×&6Ñ&6Ó&8Ó!9Ó:Ñˆ
Ü& u¬d×3Ò3��a–¸ˆà× Ñ ×4Ñ4×9Ñ9ˆØ×&Ñ&×:Ñ:×EÑEˆ
ä#:¸5ÈdÓ#SÑ ˆ˜&à/2°jÔ/@ˆz× Ñ   aÓ(¨Ö+Àbˆß(,�t�f˜GÑ$°(ˆ	ß(,�d�V˜8Ñ$°)ˆà—[‘[ˆ
àØ—y‘y Ó,Ø—i‘i 
Ó+ð
ð 	
r   r‰   N)	rD   rE   rF   rG   rH   r&   r¡   rI   rJ   rK   s   @r   r‡   r‡   ñ   s   ø‡ € ñò)÷
ö 
r   r‡   c                   ó\   a € ] tR tRt o RtR tR
V 3R lR llt]V 3R lR l4       tR	t	V t
R# )ÚMetalDequantizei  zº
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                ó   € Wn         R # r;   r‰   r‹   s   &&r   r&   ÚMetalDequantize.__init__  r�   r   Nc                ó:   <€ V ^8„  d   QhRS[ RS[R,          RS[ /# )r   r�   Úfull_layer_nameNr8   )r�   rP   )r   r   s   "€r   r   ÚMetalDequantize.__annotate__  s'   ø€ ÷ 9ñ 9¡$ð 9¹¸t½ð 9ÑY]ñ 9r   c                óh  € V P                   P                  P                  pV P                   P                  P                  p\	        V4      ^8  d   W!R,          /# VR,          ^ ,          pVR,          ^ ,          pVR,          ^ ,          p\        WgW…V4      p	W)P                  VP                  4      /# )r   zweight$r-   r.   )rŠ   r[   r   r   Úlenr…   r?   r"   )
r0   r�   r¨   rš   r   r   Ú	quantizedr-   r.   r„   s
   &&&,      r   r¡   ÚMetalDequantize.convert  sœ   € Ø× Ñ ×4Ñ4×9Ñ9ˆØ×&Ñ&×:Ñ:×EÑEˆ
äˆz‹?˜QÔØ#°	Õ%:Ð;Ð;à˜yÕ)¨!Õ,ˆ	Ø˜HÕ% aÕ(ˆØ˜YÕ'¨Õ*ˆä)¨)¸WÐRVÓWˆØ§¡¨&¯,©,Ó!7Ð8Ð8r   c                ó   <€ V ^8„  d   QhRR/# )r   r8   r   rR   )r   r   s   "€r   r   r©   +  s   ø€ ÷ ñ ˜Oñ r   c                ó   € \        4       # r;   )r   )r0   s   &r   Ú
reverse_opÚMetalDequantize.reverse_op*  s
   € ä‹}Ðr   r‰   r;   )rD   rE   rF   rG   rH   r&   r¡   Úpropertyr°   rI   rJ   rK   s   @r   r¤   r¤     s-   ø‡ € ñò)÷9ò 9ð ÷ó ör   r¤   )NNF)rH   Úcore_model_loadingr   r   Úquantizers.quantizers_utilsr   Úutilsr   r   r'   Útorch.nnr$   Ú
get_loggerrD   rX   r	   r   rV   r   ra   r   r…   r‡   r¤   rR   r   r   Ú<module>r¸      s}   ðñ÷* <Ý ?ß /ñ ×ÒÛÝð 
×	Ò	˜HÓ	%€à€òô,;�"—)‘)ô ;÷|/õd5õBô.
�Mô 
ô@�mö r   