§
    OŠtjÙ-  ã                   óF   — d Z ddlZ G d„ d¦  «        Z G d„ d¦  «        ZdS )aN  
Helper classes for working with low precision floating point types that
align with the opencompute (OCP) microscaling (MX) specification.
  * MXFP4Tensor: 4-bit E2M1 floating point data
  * MXScaleTensor: 8-bit E8M0 floating point data
Reference: https://www.opencompute.org/documents/ocp-microscaling-formats-mx-v1-0-spec-final-pdf
é    Nc                   ó4   — e Zd Zdd„Zd„ Zd„ Zd„ Zd„ Zd„ ZdS )	ÚMXFP4TensorNc                 ó  — || _         |�Lt          |t          j        ¦  «        s
J d¦   «         ‚|j         | _         |                      |¦  «        | _        dS |�!t          |t          ¦  «        r|n|f| _        dS t          d¦  «        ‚)at  
        Tensor class for working with four bit E2M1 floating point data as defined by the
        opencompute microscaling specification.


        Parameters:
        - data: A torch tensor of float32 numbers to convert to fp4e2m1 microscaling format.
        - size: The size of the tensor to create.
        - device: The device on which to create the tensor.
        Nú%Parameter data must be a torch tensorú.Either parameter data or size must be provided©	ÚdeviceÚ
isinstanceÚtorchÚTensorÚ_from_floatÚdataÚtupleÚsizeÚ
ValueError©Úselfr   r   r	   s       úO/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/triton/tools/mxfp.pyÚ__init__zMXFP4Tensor.__init__   s‰   € ð ˆŒØÐÝ˜d¥E¤LÑ1Ô1ÐZÐZÐ3ZÑZÔZÐ1Øœ+ˆDŒKØ×(Ò(¨Ñ.Ô.ˆDŒIˆIˆIØÐÝ *¨4µÑ 7Ô 7ÐE˜˜¸d¸XˆDŒIˆIˆIåÐMÑNÔNÐNó    c                 ót  — t          j        dd| j        t           j        | j        ¬¦  «        }t          j        dd| j        t           j        | j        ¬¦  «        }t          j        dd| j        t           j        | j        ¬¦  «        }|dz  |dz  z  |z                       t           j        ¦  «        | _        | S )Nr   é   ©r   Údtyper	   é   é   é   )r   Úrandintr   Úuint8r	   Útyper   )r   ÚSÚEÚMs       r   ÚrandomzMXFP4Tensor.random#   s�   € ÝŒM˜!˜Q T¤Yµe´kÈ$Ì+ÐVÑVÔVˆÝŒM˜!˜Q T¤Yµe´kÈ$Ì+ÐVÑVÔVˆÝŒM˜!˜Q T¤Yµe´kÈ$Ì+ÐVÑVÔVˆà˜1‘f  a¡Ñ(¨1Ñ,×2Ò2µ5´;Ñ?Ô?ˆŒ	Øˆr   c                 óÖ  — |t           j        k    s
J d¦   «         ‚| j        }|dz	  dz                       |¦  «        }|dz	  dz                       |¦  «        }|dz                       |¦  «        }t          j        |¦  «        }|dk    |dk    z  }| }|                     ¦   «         r�||         }	||         }
||         }t          j        d|	¦  «        }t          j        |
dk    |
|
dz
  ¦  «        }t          j        |
dk    |dz  d|dz  z   ¦  «        }|t          j        d|¦  «        z  |z  }|||<   |||dk    z  xx         dz  cc<   |                     t           j        ¦  «        S )	zŠ
        Convert fp4e2m1 data to float32.

        Returns:
        - A torch tensor of type dtype representing the fp4e2m1 data.
        zCCurrently only float32 is supported for fp4e2m1 to float conversionr   r   r   éÿÿÿÿç      à?ç      ð?r   )r   Úfloat32r   r    Ú
zeros_likeÚanyÚpowÚwhere)r   r   r   r!   r"   r#   ÚvalueÚis_zeroÚnon_zero_maskÚS_nzÚE_nzÚM_nzÚsignÚexponentÚmantissaÚvalue_nzs                   r   ÚtozMXFP4Tensor.to+   sƒ  € ð �œÒ%Ð%Ð%Ð'lÑ%Ô%Ð%àŒyˆØ�a‰i˜3Ñ×$Ò$ UÑ+Ô+ˆØ�a‰i˜3Ñ×$Ò$ UÑ+Ô+ˆØ�C‰Z×Ò˜eÑ$Ô$ˆõ Ô  Ñ#Ô#ˆØ˜’6˜a 1šfÑ%ˆØ ˜ˆØ×ÒÑÔð 	,Ø�]Ô#ˆDØ�]Ô#ˆDØ�]Ô#ˆDå”9˜R Ñ&Ô&ˆDå”{ 4¨1¢9¨d°D¸1±HÑ=Ô=ˆHÝ”{ 4¨1¢9¨d°S©j¸#ÀÀsÁ
Ñ:JÑKÔKˆHØ�eœi¨¨8Ñ4Ô4Ñ4°xÑ?ˆHà#+ˆE�-Ñ ð 	ˆg˜˜ašÑ Ð!Ð!Ô! RÑ'Ð!Ð!Ñ!Ø�zŠz�%œ-Ñ(Ô(Ð(r   c                 ó‚  — t          j        |¦  «                             t           j        ¦  «        }t          j        |¦  «        }|dk    }t          j        |¦  «        t          j        |¦  «        z  }t          j        g d¢t           j        | j        ¬¦  «        }t          j        ddgt           j        | j        ¬¦  «        }g }g }	g }
|D ]Ç}|dk    rTd}|D ]N}|dz  }|d|z  z  }| 	                    |¦  «         |	 	                    |¦  «         |
 	                    |¦  «         ŒOŒ\| 
                    ¦   «         dz
  }|D ]Q}d|dz  z   }|d|z  z  }| 	                    |¦  «         |	 	                    |¦  «         |
 	                    |¦  «         ŒRŒÈt          j        |t           j        | j        ¬¦  «        }t          j        |	t           j        | j        ¬¦  «        }	t          j        |
t           j        | j        ¬¦  «        }
|                     d¦  «        }|j        d         }|                     d¦  «        }|                     ¦   «          
                    ¦   «         }|||                     d¦  «        <   t          j        ||                     d¦  «        z
  ¦  «        }t          j        |dd	¬
¦  «        \  }}||k    }|                     ¦   «         dk    rT|
                     d¦  «                             |d¦  «        }|dk                         t           j        ¦  «        }||dz  z
  }t          j        |d¬¦  «        }|	|         }|
|         }|                     |j        ¦  «        }|                     |j        ¦  «        }d||<   d||<   |dz  |dz  z  |z                       t           j        ¦  «        S )a5  
        Convert float32 numbers to mxf4 e2m1 format.
        * No encodings are reserved for Inf or NaN in mxf4.
        * Conversion from float supports roundTiesToEven rounding mode.
        * If a value exceeds the mxf4 representable range after rounding,
          clamps to the maximum mxf4 magnitude, preserving the sign.
        * If a value has magnitude less than the minimum subnormal magnitude
          in mxf4 after rounding, converts to zero.

        Parameters:
        - values: A torch tensor of float32 numbers to convert to fp4 format.
        r   )r   r   r   r   ©r   r	   r   r'   r   r(   r&   T)ÚdimÚkeepdimg�íµ ÷Æ°>©r;   r   )r   Úsignbitr    r   ÚabsÚisnanÚisinfÚtensorr	   ÚappendÚitemr)   ÚviewÚshapeÚ	unsqueezeÚmaxÚminÚsumÚexpandÚint32Úargmin)r   Úvaluesr!   Ú
abs_valuesr/   Ú
is_invalidÚE_bitsÚM_bitsÚcandidate_valuesÚcandidate_EÚcandidate_Mr"   r5   r#   Úsignificandr.   Ú
candidatesÚabs_values_flatÚNÚabs_values_expandedÚmax_candidate_valueÚerrorsÚ
min_errorsÚ_Úis_tieÚM_bits_expandedÚtie_breakerÚbest_indicesÚ
E_selectedÚ
M_selecteds                                 r   r   zMXFP4Tensor._from_floatN   sœ  € õ ŒM˜&Ñ!Ô!×&Ò&¥u¤{Ñ3Ô3ˆÝ”Y˜vÑ&Ô&ˆ
à ’?ˆÝ”[ Ñ(Ô(­5¬;°vÑ+>Ô+>Ñ>ˆ
õ
 ”˜l˜l˜lµ%´+ÀdÄkÐRÑRÔRˆÝ”˜q !˜f­E¬KÀÄÐLÑLÔLˆàÐØˆØˆàð 	*ð 	*ˆAØ�AŠvˆvà�Øð *ð *�AØ"# c¡'�KØ'¨1¨h©;Ñ7�EØ$×+Ò+¨EÑ2Ô2Ð2Ø×&Ò& qÑ)Ô)Ð)Ø×&Ò& qÑ)Ô)Ð)Ð)ð*ð Ÿ6š6™8œ8 a™<�Øð *ð *�AØ"%¨¨C©¡-�KØ'¨1¨h©;Ñ7�EØ$×+Ò+¨EÑ2Ô2Ð2Ø×&Ò& qÑ)Ô)Ð)Ø×&Ò& qÑ)Ô)Ð)Ð)ð*õ ”\Ð"2½%¼-ÐPTÔP[Ð\Ñ\Ô\ˆ
Ý”l ;µe´kÈ$Ì+ÐVÑVÔVˆÝ”l ;µe´kÈ$Ì+ÐVÑVÔVˆà$Ÿ/š/¨"Ñ-Ô-ˆØÔ! !Ô$ˆØ-×7Ò7¸Ñ:Ô:Ðð )ŸnšnÑ.Ô.×3Ò3Ñ5Ô5ÐØ/Bˆ˜
Ÿš¨Ñ+Ô+Ñ,õ ”Ð.°×1EÒ1EÀaÑ1HÔ1HÑHÑIÔIˆõ
 œ	 &¨a¸Ð>Ñ>Ô>‰ˆ
�AØ˜JÒ&ˆà�:Š:‰<Œ<˜!ÒÐØ)×3Ò3°AÑ6Ô6×=Ò=¸aÀÑDÔDˆOØ*¨aÒ/×5Ò5µe´kÑBÔBˆKà˜{¨TÑ1Ñ2ˆFå”| F°Ð2Ñ2Ô2ˆà  Ô.ˆ
Ø  Ô.ˆ
Ø�OŠO˜JÔ,Ñ-Ô-ˆØ�OŠO˜JÔ,Ñ-Ô-ˆàˆˆ'‰
Øˆˆ'‰
à�a‘˜A ™FÑ# aÑ'×-Ò-­e¬kÑ:Ô:Ð:r   c                 ó$  — | j         }d|cxk    r|j        k     sn J d¦   «         ‚|                     |¦  «        }|dz   dz  }|dz  dk    rNdgd|j        z  z  }|j        |z
  dz
  dz  dz   }d||<   t          j        j                             ||dd¬¦  «        }t          |j        ¦  «        }|||<   | 	                    |dz   d¦  «          |j
        |Ž }|                     |dz   d¦  «        }|                     |dz   d¦  «        }	|	dz  |z  }
|
S )a  
        Packs two e2m1 elements into a single uint8 along the specified dimension.

        Parameters:
        - dim: The dimension along which to pack the elements.

        Returns:
        - A torch tensor of dtype uint8 with two e2m1 elements packed into one uint8.
        r   zHThe dimension to pack along is not within the range of tensor dimensionsr   r   Úconstant)Úmoder.   r   )r   Úndimr   r   ÚnnÚ
functionalÚpadÚlistrF   ÚinsertÚreshapeÚselect)r   r;   r   Úsize_along_dimÚnew_size_along_dimÚ	pad_sizesÚ	pad_indexÚ	new_shapeÚlowÚhighÚpackeds              r   Úto_packed_tensorzMXFP4Tensor.to_packed_tensor¦   sI  € ð ŒyˆØ�CÐ#Ð#Ò#Ð#˜$œ)Ò#Ð#Ð#Ð#Ð#ØVñ $Ô#Ð#ð Ÿš 3™œˆØ,¨qÑ0°QÑ6Ðð ˜AÑ Ò"Ð"Ø˜˜q 4¤9™}Ñ-ˆIØœ S™¨1Ñ,°Ñ1°AÑ5ˆIØ#$ˆI�iÑ Ý”8Ô&×*Ò*¨4°ÀÐSTÐ*ÑUÔUˆDå˜œÑ$Ô$ˆ	Ø+ˆ	�#‰Ø×Ò˜˜q™ !Ñ$Ô$Ð$ØˆtŒ|˜YÐ'ˆà�kŠk˜# ™' 1Ñ%Ô%ˆØ�{Š{˜3 ™7 AÑ&Ô&ˆØ˜!‘)˜sÑ"ˆàˆr   c                 óÀ  — |dz	  dz  }|dz  }t          j        ||f|dz   ¬¦  «        }t          |j        ¦  «        }|d|…         ||         dz  gz   ||dz   d…         z   } |j        |Ž }	||         dz  dk    rFt          d¦  «        g|	j        z  }
t          d||         ¦  «        |
|<   |	t          |
¦  «                 }	|	                     t           j	        ¦  «        S )aÅ  
        Unpacks a tensor where two fp4 elements are packed into a single uint8.

        Parameters:
        - packed_tensor: The packed tensor
        - dim: The dimension along which the tensor was packed.
        - original_shape: The shape of the original tensor before packing.

        Returns:
        - A tensor with the original data unpacked into uint8 elements containing one
          fp4e2m1 element in the least significant bits.
        r   é   r   r=   Nr   r   )
r   Ústackrl   rF   rn   Úslicerh   r   r    r   )r   Úpacked_tensorr;   Úoriginal_shaperv   ru   ÚstackedrF   rt   r   Úindicess              r   Úunpack_packed_tensorz MXFP4Tensor.unpack_packed_tensorÉ   sí   € ð  Ñ" cÑ)ˆØ˜cÑ!ˆå”+˜s D˜k¨s°Q©wÐ7Ñ7Ô7ˆõ �W”]Ñ#Ô#ˆØ˜$˜3˜$”K 5¨¤:°¡>Ð"2Ñ2°U¸3À¹7¸8¸8´_ÑDˆ	ØˆwŒ 	Ð*ˆð ˜#Ô Ñ" aÒ'Ð'Ý˜T‘{”{�m d¤iÑ/ˆGÝ   N°3Ô$7Ñ8Ô8ˆG�C‰LØ�˜g™œÔ'ˆDà�yŠy�œÑ%Ô%Ð%r   ©NNN)	Ú__name__Ú
__module__Ú__qualname__r   r$   r8   r   rx   r�   © r   r   r   r      s}   € € € € € ðOð Oð Oð Oð*ð ð ð!)ð !)ð !)ðFV;ð V;ð V;ðp!ð !ð !ðF&ð &ð &ð &ð &r   r   c                   ó*   — e Zd Zdd„Zdd„Zd„ Zd„ ZdS )ÚMXScaleTensorNc                 ó  — || _         |�Lt          |t          j        ¦  «        s
J d¦   «         ‚|j         | _         |                      |¦  «        | _        dS |�!t          |t          ¦  «        r|n|f| _        dS t          d¦  «        ‚)a6  
        Tensor class for working with microscaling E8M0 block scale factors.

        Parameters:
        - data: A torch tensor of float32 numbers to convert to fp8e8m0 microscaling format.
        - size: The size of the tensor to create.
        - device: The device on which to create the tensor.
        Nr   r   r   r   s       r   r   zMXScaleTensor.__init__ë   s‰   € ð ˆŒØÐÝ˜d¥E¤LÑ1Ô1ÐZÐZÐ3ZÑZÔZÐ1Øœ+ˆDŒKØ×(Ò(¨Ñ.Ô.ˆDŒIˆIˆIØÐÝ *¨4µÑ 7Ô 7ÐE˜˜¸d¸XˆDŒIˆIˆIåÐMÑNÔNÐNr   c                 óÔ  — d}|€dnCt          dt          t          j        t          j        |¦  «        ¦  «        ¦  «        |z   ¦  «        }|€dnQt          dt          dt          t          j        t          j        |¦  «        ¦  «        ¦  «        |z   ¦  «        ¦  «        }||k    s
J d¦   «         ‚t          j        ||dz   | j        t          j        | j	        ¬¦  «        }|| _
        | S )zp
        Generate random E8M0 data within a specified range.
        * Excludes the NaN encoding (255).
        é   Nr   éþ   z&Low must be less than or equal to highr   r   )rH   Úintr   Úlog2rB   rI   r   r   r   r	   r   )r   ru   rv   ÚbiasÚmin_exponentÚmax_exponentr"   s          r   r$   zMXScaleTensor.randomþ   sÖ   € ð
 ˆà˜K�q�q­S°µC½¼
Å5Ä<ÐPSÑCTÔCTÑ8UÔ8UÑ4VÔ4VÐY]Ñ4]Ñ-^Ô-^ˆØ"˜l�s�sµ°C½¸QÅÅEÄJÍuÌ|Ð\`ÑOaÔOaÑDbÔDbÑ@cÔ@cÐfjÑ@jÑ9kÔ9kÑ0lÔ0lˆØ˜|Ò+Ð+Ð+Ð-UÑ+Ô+Ð+åŒM˜,¨°qÑ(8¸t¼yÕPUÔP[ÐdhÔdoÐpÑpÔpˆØˆŒ	Øˆr   c                 ó$  — |t           j        k    s
J d¦   «         ‚| j                             |¦  «        }|dk    }|                     ¦   «         }d||<   |dz
  }t          j        d|¦  «        }t           j        ||<   |                     |¦  «        S )NzBCurrently only float32 is supported for f8e8m0 to float conversionéÿ   r   r‹   g       @)r   r)   r   r    Úcloner,   Únan)r   r   r   Úis_nanÚe_biasedÚer.   s          r   r8   zMXScaleTensor.to  sˆ   € Ø�œÒ%Ð%Ð%Ð'kÑ%Ô%Ð%ØŒy�~Š~˜eÑ$Ô$ˆØ˜#’+ˆØ—:’:‘<”<ˆØˆ�ÑØ�s‰NˆÝ”	˜#˜qÑ!Ô!ˆÝœ	ˆˆf‰Ø�zŠz˜%Ñ Ô Ð r   c                 óÔ  — t          j        |t           j        | j        ¬¦  «        }t          j        |¦  «        t          j        |¦  «        z  |dk    z  }d||<   ||          }t          j        t          j        |¦  «        ¦  «        }|dz   }|                     t           j	        ¦  «        }t          j
        |dd¦  «        }|                     t           j        ¦  «        || <   |S )aO  
        Convert float32 numbers to E8M0 format.
        * Values <= 0, NaNs, and Infs are converted to the NaN encoding (255).
        * Positive values are converted by computing the floor of log2(value) to get the exponent.

        Parameters:
        - values: A torch tensor of float32 numbers to convert to E8M0 format.
        r:   r   r“   r‹   rŒ   )r   Ú
empty_liker   r	   r@   rA   ÚfloorrŽ   r    rL   Úclamp)	r   rN   ÚresultrP   Úvalid_valuesr˜   r—   Úe_biased_intÚe_biased_clampeds	            r   r   zMXScaleTensor._from_float  sÆ   € õ Ô! &µ´ÀDÄKÐPÑPÔPˆå”[ Ñ(Ô(­5¬;°vÑ+>Ô+>Ñ>À&ÈAÂ+ÑNˆ
Ø ˆˆzÑà˜z˜kÔ*ˆÝŒK�œ
 <Ñ0Ô0Ñ1Ô1ˆØ�s‘7ˆØ—}’}¥U¤[Ñ1Ô1ˆÝ œ; |°Q¸Ñ<Ô<ÐØ.×3Ò3µE´KÑ@Ô@ˆ�
ˆ{Ñàˆr   r‚   )NN)rƒ   r„   r…   r   r$   r8   r   r†   r   r   rˆ   rˆ   é   s^   € € € € € ðOð Oð Oð Oð&ð ð ð ð	!ð 	!ð 	!ðð ð ð ð r   rˆ   )Ú__doc__r   r   rˆ   r†   r   r   ú<module>r¢      s‡   ððð ð €€€ðZ&ð Z&ð Z&ð Z&ð Z&ñ Z&ô Z&ð Z&ðzDð Dð Dð Dð Dñ Dô Dð Dð Dð Dr   