Ë
    7^(hÙ-  ã                   ó8   — d Z ddlZ G d„ d«      Z G d„ d«      Zy)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                   ó2   — e Zd Zdd„Zd„ Zd„ Zd„ Zd„ Zd„ Zy)	ÚMXFP4TensorNc                 ó  — || _         |�It        |t        j                  «      sJ d«       ‚|j                   | _         | j	                  |«      | _        y|�!t        |t        «      r|| _        y|f| _        y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/skyplay_api_hub/venv/lib/python3.12/site-packages/triton/tools/mxfp.pyÚ__init__zMXFP4Tensor.__init__   sr   € ð ˆŒØÐÜ˜d¤E§L¡LÔ1ÐZÐ3ZÓZÐ1ØŸ+™+ˆDŒKØ×(Ñ(¨Ó.ˆD�IØÐÜ *¨4´Ô 7˜ˆD�I¸d¸XˆD�IäÐMÓNÐNó    c                 óÎ  — 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  j                  t         j                  «      | _        | S )Nr   é   ©r   Údtyper	   é   é   é   )r   Úrandintr   Úuint8r	   Útyper   )r   ÚSÚEÚMs       r   ÚrandomzMXFP4Tensor.random#   s•   € Ü�M‰M˜!˜Q T§Y¡Y´e·k±kÈ$Ï+É+ÔVˆÜ�M‰M˜!˜Q T§Y¡Y´e·k±kÈ$Ï+É+ÔVˆÜ�M‰M˜!˜Q T§Y¡Y´e·k±kÈ$Ï+É+ÔVˆà˜1‘f  a¡Ñ(¨1Ñ,×2Ñ2´5·;±;Ó?ˆŒ	Øˆr   c                 ó¨  — |t         j                  k(  sJ d«       ‚| j                  }|dz	  dz  j                  |«      }|dz	  dz  j                  |«      }|dz  j                  |«      }t        j                  |«      }|dk(  |dk(  z  }| }|j                  «       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<   |j                  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Ð'lÓlÐ%à�y‰yˆØ�a‰i˜3Ñ×$Ñ$ UÓ+ˆØ�a‰i˜3Ñ×$Ñ$ UÓ+ˆØ�C‰Z×Ñ˜eÓ$ˆô × Ñ  Ó#ˆØ˜‘6˜a 1™fÑ%ˆØ ˜ˆØ×ÑÔØ�]Ñ#ˆDØ�]Ñ#ˆDØ�]Ñ#ˆDä—9‘9˜R Ó&ˆDä—{‘{ 4¨1¡9¨d°D¸1±HÓ=ˆHÜ—{‘{ 4¨1¡9¨d°S©j¸#ÀÀsÁ
Ñ:JÓKˆHØœeŸi™i¨¨8Ó4Ñ4°xÑ?ˆHà#+ˆE�-Ñ ð 	ˆg˜˜a™Ñ Ó! RÑ'Ó!Ø�z‰zœ%Ÿ-™-Ó(Ð(r   c                 ór  — t        j                  |«      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(  rJd}|D ]B  }|dz  }|d|z  z  }|j                  |«       |	j                  |«       |
j                  |«       ŒD ŒR|j                  «       dz
  }|D ]E  }d|dz  z   }|d|z  z  }|j                  |«       |	j                  |«       |
j                  |«       ŒG Œ° t        j                  |t         j                  | j                  ¬«      }t        j                  |	t         j                  | j                  ¬«      }	t        j                  |
t         j                  | j                  ¬«      }
|j                  d«      }|j                  d   }|j                  d«      }|j                  «       j                  «       }|||j                  d«      <   t        j                  ||j                  d«      z
  «      }t        j                   |dd	¬
«      \  }}||k(  }|j#                  «       dkD  rK|
j                  d«      j%                  |d«      }|dk(  j                  t         j&                  «      }||dz  z
  }t        j(                  |d¬«      }|	|   }|
|   }|j                  |j                  «      }|j                  |j                  «      }d||<   d||<   |dz  |dz  z  |z  j                  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   s7  € ô �M‰M˜&Ó!×&Ñ&¤u§{¡{Ó3ˆÜ—Y‘Y˜vÓ&ˆ
à ‘?ˆÜ—[‘[ Ó(¬5¯;©;°vÓ+>Ñ>ˆ
ô
 —‘šl´%·+±+ÀdÇkÁkÔRˆÜ—‘˜q !˜f¬E¯K©KÀÇÁÔLˆàÐØˆØˆàò 	*ˆAØ�AŠvà�Øò *�AØ"# c¡'�KØ'¨1¨h©;Ñ7�EØ$×+Ñ+¨EÔ2Ø×&Ñ& qÔ)Ø×&Ñ& qÕ)ñ*ð Ÿ6™6›8 a™<�Øò *�AØ"%¨¨C©¡-�KØ'¨1¨h©;Ñ7�EØ$×+Ñ+¨EÔ2Ø×&Ñ& qÔ)Ø×&Ñ& qÕ)ñ*ð	*ô( —\‘\Ð"2¼%¿-¹-ÐPT×P[ÑP[Ô\ˆ
Ü—l‘l ;´e·k±kÈ$Ï+É+ÔVˆÜ—l‘l ;´e·k±kÈ$Ï+É+ÔVˆà$Ÿ/™/¨"Ó-ˆØ×!Ñ! !Ñ$ˆØ-×7Ñ7¸Ó:Ðð )Ÿn™nÓ.×3Ñ3Ó5ÐØ/Bˆ˜
Ÿ™¨Ó+Ñ,ô —‘Ð.°×1EÑ1EÀaÓ1HÑHÓIˆô
 Ÿ	™	 &¨a¸Ô>‰ˆ
�AØ˜JÑ&ˆà�:‰:‹<˜!ÒØ)×3Ñ3°AÓ6×=Ñ=¸aÀÓDˆOØ*¨aÑ/×5Ñ5´e·k±kÓBˆKà˜{¨TÑ1Ñ2ˆFä—|‘| F°Ô2ˆà  Ñ.ˆ
Ø  Ñ.ˆ
Ø�O‰O˜J×,Ñ,Ó-ˆØ�O‰O˜J×,Ñ,Ó-ˆàˆˆ'‰
Øˆˆ'‰
à�a‘˜A ™FÑ# aÑ'×-Ñ-¬e¯k©kÓ:Ð:r   c                 óB  — | j                   }d|cxk  r|j                  k  sJ d«       ‚ J d«       ‚|j                  |«      }|dz   dz  }|dz  dk7  r]dgd|j                  z  z  }|j                  |z
  dz
  dz  dz   }d||<   t        j                  j
                  j                  ||dd¬«      }t        |j                  «      }|||<   |j                  |dz   d«        |j                  |Ž }|j                  |dz   d«      }|j                  |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¦   sD  € ð �y‰yˆØ�CÔ#˜$Ÿ)™)Ò#ð 	WØVó	WÑ#ð 	WØVó	WÐ#ð Ÿ™ 3›ˆØ,¨qÑ0°QÑ6Ðð ˜AÑ Ò"Ø˜˜q 4§9¡9™}Ñ-ˆIØŸ™ S™¨1Ñ,°Ñ1°AÑ5ˆIØ#$ˆI�iÑ Ü—8‘8×&Ñ&×*Ñ*¨4°ÀÐSTÐ*Ó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k7  r9t        d«      g|	j                  z  }
t        d||   «      |
|<   |	t        |
«         }	|	j                  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ˆô �W—]‘]Ó#ˆØ˜$˜3�K 5¨¡:°¡>Ð"2Ñ2°U¸3À¹7¸8°_ÑDˆ	Øˆw�‰ 	Ð*ˆð ˜#Ñ Ñ" aÒ'Ü˜T“{�m d§i¡iÑ/ˆGÜ   N°3Ñ$7Ó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ò*ò!)òFV;òp!óF&r   r   c                   ó(   — e Zd Zdd„Zdd„Zd„ Zd„ Zy)ÚMXScaleTensorNc                 ó  — || _         |�It        |t        j                  «      sJ d«       ‚|j                   | _         | j	                  |«      | _        y|�!t        |t        «      r|| _        y|f| _        y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__ë   sr   € ð ˆŒØÐÜ˜d¤E§L¡LÔ1ÐZÐ3ZÓZÐ1ØŸ+™+ˆDŒKØ×(Ñ(¨Ó.ˆD�IØÐÜ *¨4´Ô 7˜ˆD�I¸d¸XˆD�IäÐMÓNÐNr   c                 óÊ  — d}|€dn=t        dt        t        j                  t        j                  |«      «      «      |z   «      }|€dnGt        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).
        é   r   éþ   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¬S°´C¼¿
¹
Ä5Ç<Á<ÐPSÓCTÓ8UÓ4VÐY]Ñ4]Ó-^ˆØ"˜l‘s´°C¼¸QÄÄEÇJÁJÌuÏ|É|Ð\`ÓOaÓDbÓ@cÐfjÑ@jÓ9kÓ0lˆØ˜|Ò+ÐUÐ-UÓUÐ+ä�M‰M˜,¨°qÑ(8¸t¿y¹yÔPU×P[ÑP[Ðdh×doÑdoÔpˆØˆŒ	Øˆr   c                 ó  — |t         j                  k(  sJ d«       ‚| j                  j                  |«      }|dk(  }|j	                  «       }d||<   |dz
  }t        j
                  d|«      }t         j                  ||<   |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Ð'kÓkÐ%Ø�y‰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   }|j                  t         j                  «      }t        j                  |dd«      }|j                  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ÁKÔPˆä—[‘[ Ó(¬5¯;©;°vÓ+>Ñ>À&ÈAÁ+ÑNˆ
Ø ˆˆzÑà˜z˜kÑ*ˆÜ�K‰KœŸ
™
 <Ó0Ó1ˆØ�s‘7ˆØ—}‘}¤U§[¡[Ó1ˆÜ Ÿ;™; |°Q¸Ó<ÐØ.×3Ñ3´E·K±KÓ@ˆ�
ˆ{Ñàˆr   r‚   )NN)rƒ   r„   r…   r   r$   r8   r   r†   r   r   rˆ   rˆ   é   s   „ óOó&ò	!ór   rˆ   )Ú__doc__r   r   rˆ   r†   r   r   ú<module>r¢      s(   ðñó ÷Z&ñ Z&÷zDò Dr   