§
    ‚Štj)  ã            
       óh  — 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dS )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d¬¦  «        a n&# t          $ r}t	          d|› d�¦  «        |‚d}~ww xY wt           S )z>Lazily load the quantization-mlx kernel from Hugging Face Hub.Né   )Ú
get_kernelz0kernels-community/mlx-quantization-metal-kernels)Úversionz9Failed 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     új/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/transformers/integrations/metal_quantization.pyÚ_get_metal_kernelr   3   s�   € õ Ðð		Ø/Ð/Ð/Ð/Ð/Ð/à&˜JÐ'YÐcdÐeÑeÔeˆMˆMøÝð 	ð 	ð 	Ýð?ÈAð ?ð ?ð ?ñô ð ð	øøøøð	øøøõ Ðs   ‰ �
A §;»A c                   óf   — 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
dS )Ú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                             | ¦  «         || _        || _        || _        || _        d|z  }||z  }||z  }	|t          j        k    r;t          j	        t          j
        ||t          j        ¬¦  «        d¬¦  «        | _        n0t          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
        |¦  «        ¦  «        | _        d S |                      dd ¦  «         d S )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   sc  € õ 	Œ	×Ò˜4Ñ Ô Ð à&ˆÔØ(ˆÔØˆŒ	Ø$ˆŒà˜d™
ˆØ -Ñ/ˆØ *Ñ,ˆà•E”LÒ Ð Ýœ,¥u¤{°<ÀÕQVÔQ]Ð'^Ñ'^Ô'^ÐnsÐtÑtÔtˆDŒKˆKåœ,¥u¤{°<ÀÐTYÐ'ZÑ'ZÔ'ZÐjoÐpÑpÔpˆDŒKà(-µ´Ò(=Ð(=•u”}�}À4ˆÝ”l¥5¤;¨|¸XÈ\Ð#ZÑ#ZÔ#ZÐjoÐpÑpÔpˆŒÝ”|¥E¤K°¸hÈlÐ$[Ñ$[Ô$[ÐkpÐqÑqÔqˆŒàð 	2Ýœ¥U¤[°Ñ%>Ô%>Ñ?Ô?ˆDŒIˆIˆIà×#Ò# F¨DÑ1Ô1Ð1Ð1Ð1ó    ÚinputÚreturnc                 ó”  — | j         j        t          j        k    r+t          j                             || j         | j        ¦  «        S t          ¦   «         }| 	                    || 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-   r3   ÚkernelÚoutputs       r   ÚforwardzMetalLinear.forwards   s¥   € ØŒ;Ô¥¤Ò,Ð,Ý”=×'Ò'¨¨t¬{¸D¼IÑFÔFÐFå"Ñ$Ô$ˆà×$Ò$ØØŒKØŒK�NŠN˜5œ;Ñ'Ô'ØŒL�OŠO˜EœKÑ(Ô(ØŒOØŒIñ
ô 
ˆð Œ9Ð Ø˜dœiÑ'ˆFØˆr2   N)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r$   r%   ÚintÚboolr#   ÚTensorr=   © r2   r   r   r   I   sž   € € € € € ðð ð ØŒlØØð 2ð  2àð 2ð ð 2ð ð	 2ð ð 2ð ð 2ð  2ð  2ð  2ðD˜Uœ\ð ¨e¬lð ð ð ð ð ð r2   r   FÚmodules_to_not_convertÚpre_quantizedc           
      ó€  — |j         r| S |j        }|j        }d}|                      ¦   «         D ]s\  }}t	          ||¦  «        sŒt          |t          j        ¦  «        rC|ri nddi}	t          d|j	        |j
        |j        du||dœ|	¤Ž}
|                      ||
¦  «         d}Œt|st                               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.rE   )Ú
dequantizer   r   Únamed_modulesr   Ú
isinstancer!   ÚLinearr   r   r   r   Úset_submoduleÚloggerÚwarning)ÚmodelrF   Úquantization_configrG   r   r   Úhas_been_replacedÚmodule_nameÚmoduleÚmodule_kwargsÚ
new_modules              r   Úreplace_with_metal_linearrW   ‡   s  € ð Ô%ð ØˆàÔ#€DØ$Ô/€JàÐà$×2Ò2Ñ4Ô4ð %ð %Ñˆ�VÝ$ [Ð2HÑIÔIð 	Øå�f�bœiÑ(Ô(ð 	%Ø"/ÐD˜B˜B°g¸t°_ˆMÝ$ð Ø"Ô.Ø#Ô0Ø”[¨Ð,ØØ%ðð ð  ðð ˆJð ×Ò ¨ZÑ8Ô8Ð8Ø $Ðøàð 
Ý�Šð;ñ	
ô 	
ð 	
ð
 €Lr2   r(   r   r   c                 ó2  — | j         \  }}d|z  }d|z  dz
  }||z  }|                      ¦   «                              |||¦  «        }|                     d¬¦  «        j        }	|                     d¬¦  «        j        }
|
|	z
  |z                       d¬¦  «        }|	}||                     d¦  «        z
  |                     d¦  «        z  }|                     ¦   «                              d|¦  «         	                    t          j        ¦  «                             ||¦  «        }||z  }t          j        ||t          j        | j        ¬¦  «        }t          |¦  «        D ]}||d	d	…|d	|…f         ||z  z  z  }Œ| 	                    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Úreshaper[   ÚvaluesÚmaxÚclampÚ	unsqueezeÚroundr:   r$   Úint32r'   r]   Úranger%   )r(   r   r   ÚNÚKr.   Úmax_valr0   Ú	w_groupedÚw_minÚw_maxr*   ÚbiasesÚw_intr/   Úw_packedÚis                    r   Ú_affine_quantize_tensorrr   ¹   sŠ  € ð Œ<�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Ô8€FØ€Fà˜×)Ò)¨"Ñ-Ô-Ñ-°×1AÒ1AÀ"Ñ1EÔ1EÑE€EØ�KŠK‰MŒM×Ò  7Ñ+Ô+×.Ò.­u¬{Ñ;Ô;×CÒCÀAÀqÑIÔI€Eð �MÑ!€HÝŒ{˜1˜h­e¬kÀ&Ä-ÐPÑPÔP€HÝ�=Ñ!Ô!ð =ð =ˆØ�E˜!˜!˜!˜QÐ- Ð-Ð-Ô.°4¸!±8Ñ<Ñ<ˆˆà�;Š;•u”|Ñ$Ô$ f¨fÐ4Ð4r2   rp   r*   rn   c                 óR  — | j         d         }d|z  }d|z  dz
  }| j         d         |z  }|                      t          j        ¦  «        }	t          j        ||t          j        | j        ¬¦  «        }
t          |¦  «        D ])}|	||z  z	  |z                       ¦   «         |
dd…|d|…f<   Œ*|
 	                    |d|¦  «        }||                     ¦   «          
                    d¦  «        z  |                     ¦   «          
                    d¦  «        z   }| 	                    ||¦  «        S )zv
    Dequantize a packed uint32 weight ``[N, K_packed]`` back to float.

    Returns a ``[N, K]`` float32 tensor.
    r   r   r
   r\   NrY   )r^   r:   r$   rf   r'   r)   r]   rg   r_   r`   rd   )rp   r*   rn   r   r   rh   r.   rj   ri   Ú
w_packed_iÚw_flatrq   rk   Úw_deqs                 r   Ú_affine_dequantize_tensorrw   Ú   s  € ð 	Œ�qÔ€AØ˜$‘J€MØ�D‰y˜A‰o€GØŒ�qÔ˜MÑ)€Aà—’�Uœ[Ñ)Ô)€JÝŒ[˜˜A¥U¤]¸8¼?ÐKÑKÔK€FÝ�=Ñ!Ô!ð Uð UˆØ(2°t¸a±xÑ(@ÀGÑ'K×&RÒ&RÑ&TÔ&Tˆˆqˆqˆq�!Ð"�]Ð"Ð"Ñ#Ð#à—’˜q " jÑ1Ô1€IØ˜Ÿš™œ×0Ò0°Ñ4Ô4Ñ4°v·|²|±~´~×7OÒ7OÐPRÑ7SÔ7SÑS€EØ�=Š=˜˜AÑÔÐr2   c                   ó(   — e Zd ZdZd„ Zdedefd„ZdS )Ú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                 ó   — || _         d S r6   ©Úhf_quantizer©r-   r|   s     r   r#   zMetalQuantize.__init__ù   ó   € Ø(ˆÔÐÐr2   Ú
input_dictr4   c                 óâ  — t          t          |                     ¦   «         ¦  «        ¦  «        \  }}t          |t          ¦  «        r|d         n|}| j        j        j        }| j        j        j        }t          |||¦  «        \  }}}	d|v r| 
                    dd¦  «        d         nd}
|
r|
› d�nd}|
r|
› d�nd}|j        }||||                     |¦  «        ||	                     |¦  «        iS )	Nr   ú.r
   Ú z.scalesr*   z.qbiasesr+   )ÚnextÚiterÚitemsrK   Úlistr|   rQ   r   r   rr   Úrsplitr   r:   )r-   r   ÚkwargsÚ
target_keyÚvaluer   r   rp   r*   rn   ÚbaseÚ	scale_keyÚbias_keyÚ
orig_dtypes                 r   ÚconvertzMetalQuantize.convertü   s	  € Ý ¥ j×&6Ò&6Ñ&8Ô&8Ñ!9Ô!9Ñ:Ô:Ñˆ
�EÝ& u­dÑ3Ô3Ð>��a”�¸ˆàÔ Ô4Ô9ˆØÔ&Ô:ÔEˆ
å#:¸5À*ÈdÑ#SÔ#SÑ ˆ�&˜&à/2°jÐ/@Ð/@ˆz× Ò   aÑ(Ô(¨Ô+Ð+ÀbˆØ(,Ð:�tÐ$Ð$Ð$Ð$°(ˆ	Ø(,Ð;�dÐ$Ð$Ð$Ð$°)ˆà”[ˆ
à˜Ø�v—y’y Ñ,Ô,Ø�f—i’i 
Ñ+Ô+ð
ð 	
r2   N)r>   r?   r@   rA   r#   Údictr�   rE   r2   r   ry   ry   ñ   sO   € € € € € ðð ð)ð )ð )ð
 $ð 
°Tð 
ð 
ð 
ð 
ð 
ð 
r2   ry   c                   óL   — e Zd ZdZd„ Zd
dededz  defd„Zedd	„¦   «         Z	dS )Ú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                 ó   — || _         d S r6   r{   r}   s     r   r#   zMetalDequantize.__init__  r~   r2   Nr   Úfull_layer_namer4   c                 ó2  — | j         j        j        }| j         j        j        }t	          |¦  «        dk     r
||d         iS |d         d         }|d         d         }|d         d         }t          |||||¦  «        }	||	                     |j        ¦  «        iS )Nr   zweight$r   r*   r+   )r|   rQ   r   r   Úlenrw   r:   r   )
r-   r   r”   rˆ   r   r   Ú	quantizedr*   r+   rv   s
             r   r�   zMetalDequantize.convert  s›   € ØÔ Ô4Ô9ˆØÔ&Ô:ÔEˆ
åˆz‰?Œ?˜QÒÐØ# Z°	Ô%:Ð;Ð;à˜yÔ)¨!Ô,ˆ	Ø˜HÔ% aÔ(ˆØ˜YÔ'¨Ô*ˆå)¨)°V¸WÀjÐRVÑWÔWˆØ §¢¨&¬,Ñ!7Ô!7Ð8Ð8r2   r   c                 ó   — t          ¦   «         S r6   )r   )r-   s    r   Ú
reverse_opzMetalDequantize.reverse_op*  s   € å‰}Œ}Ðr2   r6   )r4   r   )
r>   r?   r@   rA   r#   r�   Ústrr�   Úpropertyr™   rE   r2   r   r’   r’     s€   € € € € € ðð ð)ð )ð )ð9ð 9 $ð 9¸¸t¹ð 9ÐY]ð 9ð 9ð 9ð 9ð ðð ð ñ „Xðð ð r2   r’   )NNF)rA   Úcore_model_loadingr   r   Úquantizers.quantizers_utilsr   Úutilsr   r   r$   Útorch.nnr!   Ú
get_loggerr>   rN   r   r   rL   r   r†   rš   rC   rW   rD   rB   rr   rw   ry   r’   rE   r2   r   ú<module>r¡      sã  ððð ð* <Ð ;Ð ;Ð ;Ð ;Ð ;Ð ;Ð ;Ø ?Ð ?Ð ?Ð ?Ð ?Ð ?Ø /Ð /Ð /Ð /Ð /Ð /Ð /Ð /ð ÐÑÔð Ø€L€L€LØÐÐÐÐÐð 
ˆÔ	˜HÑ	%Ô	%€à€ðð ð ð,;ð ;ð ;ð ;ð ;�"”)ñ ;ô ;ð ;ð@ 04ØØð	/ð /à  œI¨Ñ,ð/ð ð	/ð /ð /ð /ðd5 E¤Lð 5¸cð 5Èð 5ð 5ð 5ð 5ðBØŒlðØ$)¤LðØ:?¼,ðØTWðØ_bðð ð ð ð.
ð 
ð 
ð 
ð 
�Mñ 
ô 
ð 
ð@ð ð ð ð �mñ ô ð ð ð r2   