§
    �Štj[ƒ  ã            	       óJ  — d dl Z d dlmZ d dlmZ d dlmZ d dlZd dlm	Z	m
Z
 d dlmZmZmZmZmZmZmZmZmZmZmZmZmZ g d¢Z edd	¦  «        Z G d
„ dej        ¦  «        Zdej        fdej        dededefd„Z  G d„ de¦  «        Z! G d„ de¦  «        Z"da#d„ Z$da%d„ Z&dS )é    N)Ú
namedtuple)ÚCallable)ÚAny)Ú)sparse_semi_structured_from_dense_cutlassÚ'sparse_semi_structured_to_dense_cutlass)Úfallback_dispatcherÚsemi_sparse_addmmÚsemi_sparse_cloneÚsemi_sparse_detachÚsemi_sparse_indicesÚsemi_sparse_linearÚsemi_sparse_mmÚsemi_sparse_scaled_mmÚsemi_sparse_tÚsemi_sparse_toÚsemi_sparse_to_copyÚsemi_sparse_valuesÚsemi_sparse_view)ÚSparseSemiStructuredTensorÚ!SparseSemiStructuredTensorCUTLASSÚ$SparseSemiStructuredTensorCUSPARSELTÚto_sparse_semi_structuredÚ_SEMI_STRUCTURED_SPARSE_CONFIGz=sparse_min_rows sparse_min_cols dense_min_rows dense_min_colsc                   óX  — e Zd ZU dZdZeed<   eej	        e
f         ed<   dZeed<   dZeed<   dZeed<   eed	<   eeef         ed
<   ej        dz  ed<   ej        dz  ed<   ej        dz  ed<   ej        dz  ed<   ej        dz  ed<   eed<   eed<   g d¢Ze	 	 	 d'dej        dej        dz  dej        dz  dej        dz  dej        dz  dej        dz  dededefd„¦   «         Zdefd„Zdeee         eej        eeef         f         fd„Zedeej        eeef         dej        fd„¦   «         Zej        j        Zedefd„¦   «         Z ed(d)d„¦   «         Z!edej        ddfd„¦   «         Z"d „ Z#eefdej        d!edd fd"„¦   «         Z$dd#œd$ej        d%ej        dz  dej        fd&„Z%dS )*r   a¼  
    This class implements semi-structured sparsity as a Tensor subclass.

    Semi-structured sparsity describes a sparsity pattern where n in every 2n elements are sparse,
    depending on the datatype. It is also referred to as 2:4 sparsity or fine-grained
    structured sparsity.

    There are two backends available for semi_structred sparsity, either cuSPARSELt or CUTLASS.
    This class is meant to serve as a base class for both implementations. SparseSemiStructuredCUTLASS
    and SparseSemiStructuredCUSPARSELT both inherit from this class and define three backend-specific items.
    Note that as such, this class cannot be instantiated directly.

    -`_DTYPE_SHAPE_CONSTRAINTS` - A dictionary holding backend specific dense/sparse min shape constraints
    - `def from_dense()` - backend specific compression routines
    - `def _mm()` - backend specific mm op (either torch._cslt_sparse_mm or torch._sparse_semi_structured_(mm|addmm))
    r   Ú_DEFAULT_ALG_IDÚ_DTYPE_SHAPE_CONSTRAINTSFÚ_FORCE_CUTLASSÚ_FUSE_TRANSPOSEÚ_PROTOTYPE_WARNING_SHOWNÚBACKENDÚSPARSE_DISPATCHNÚpackedÚmetaÚpacked_tÚmeta_tÚcompressed_swizzled_bitmaskÚfuse_transpose_cusparseltÚalg_id_cusparselt)r"   r#   r$   r%   r&   ÚshapeÚrequires_gradc
                 ó¼  — | j         sVt          j        dt          d¬¦  «         d| _         |                      ¦   «          t
          j                             | ¦  «         |�|}
n|�|}
nt          d¦  «        ‚t
          j	         
                    | ||
j        |
j        |
j        |	¬¦  «        }||_        ||_        ||_        ||_        ||_        ||_        ||_        |S )a0  
        Create a new instance of the tensor subclass from the compressed sparse representation.

        We have the option to create the subclass with the compressed representations of both X and X', for training.
        For inference, we only need a single representation (either X or X'), while the corresponding other set will be None.

        Depending on the backend selected, certain fields will be set to None. (CUSPARSELT vs CUTLASS)

        Args:
            shape: The shape of the original dense tensor
            packed: The compressed representation of the original dense tensor
            meta: The metadata of the original dense tensor, if it is stored separately
            packed_t: The compressed representation of the transposed original dense tensor
            meta_t: The metadata of the transposed original dense tensor, if it is stored separately
            compressed_swizzled_bitmask: The masks used by the CUTLASS backend to determine which threads should
                                         participate in the computation. Used for pointwise ops.
            fuse_transpose_cusparselt: When running with cuSPARSELt, we have the option to fuse a transposition
                                       with a matmul, which is useful in the case of 2:4 sparse training.
            alg_id_cusparselt: The algorithm id to use when using cuSPARSELT, will have effect on performance

        Returns:
            torch.Tensor: A torch.Tensor wrapper subclass.

        Raises:
            ValueError: If all of the tensor arguments are None.
        zøThe PyTorch API of SparseSemiStructuredTensor is in prototype stage and will change in the near future. Please open a Github issue for features requests and see our documentation on the torch.sparse module for further information about the project.é   ©Ú
stacklevelTNz3At least one of packed or packed_t must be provided)ÚdeviceÚdtypeÚlayoutr*   )r   ÚwarningsÚwarnÚUserWarningÚ_load_dispatch_tableÚtorchÚ_dynamoÚallow_in_graphÚ
ValueErrorÚTensorÚ_make_wrapper_subclassr/   r0   r1   r"   r#   r$   r%   r&   r'   r(   )Úclsr)   r"   r#   r$   r%   r&   r'   r(   r*   Úprevious_tensorÚtensors               úZ/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/torch/sparse/semi_structured.pyÚ__new__z"SparseSemiStructuredTensor.__new__O   s  € ðN Ô+ð 	.ÝŒMðHõ
 Øð	ñ 	ô 	ð 	ð ,0ˆCÔ(ð
 ×$Ò$Ñ&Ô&Ð&õ ŒM×(Ò(¨Ñ-Ô-Ð-àÐØ$ˆOˆOØÐ!Ø&ˆOˆOåÐRÑSÔSÐSå”×4Ò4ØØØ"Ô)Ø!Ô'Ø"Ô)Ø'ð 5ñ 
ô 
ˆð ˆŒØˆŒØ"ˆŒØˆŒØ-HˆÔ*Ø+DˆÔ(Ø#4ˆÔ Øˆó    Úreturnc                 ón   — t          | d¦  «        st          d¦  «        ‚| j        j        › d| j        › d�S )Nr)   ztensor has no shape attributez(shape=ú))ÚhasattrÚAssertionErrorÚ	__class__Ú__name__r)   )Úselfs    r?   Ú__repr__z#SparseSemiStructuredTensor.__repr__¤   sB   € Ý�t˜WÑ%Ô%ð 	BÝ Ð!@ÑAÔAÐAØ”.Ô)Ð?Ð?°$´*Ð?Ð?Ð?Ð?rA   c                 óŠ   ‡ — t          t          ˆ fd„‰ j        ¦  «        ¦  «        }‰ j        ‰ j        ‰ j        ‰ j        f}||fS )Nc                 ó(   •— t          ‰| ¦  «        d uS ©N)Úgetattr)ÚxrI   s    €r?   ú<lambda>z?SparseSemiStructuredTensor.__tensor_flatten__.<locals>.<lambda>­   s   ø€ �W T¨1Ñ-Ô-°TÐ9€ rA   )ÚlistÚfilterÚ	__slots__r)   r'   r(   r*   )rI   Úinner_tensorsÚtensor_metas   `  r?   Ú__tensor_flatten__z-SparseSemiStructuredTensor.__tensor_flatten__©   sZ   ø€ õ ÝÐ9Ð9Ð9Ð9¸4¼>ÑJÔJñ
ô 
ˆð ŒJØÔ*ØÔ"ØÔð	
ˆð ˜kÐ)Ð)rA   rU   c                 ó   — |\  }}}} | ||                      dd ¦  «        |                      dd ¦  «        |                      dd ¦  «        |                      dd ¦  «        |                      dd ¦  «        |||¬¦	  «	        S )Nr"   r#   r$   r%   r&   ©	r)   r"   r#   r$   r%   r&   r'   r(   r*   )Úget)	r<   rT   rU   Ú
outer_sizeÚouter_strider)   r'   r(   r*   s	            r?   Ú__tensor_unflatten__z/SparseSemiStructuredTensor.__tensor_unflatten__·   s¢   € ð NYÑJˆÐ(Ð*;¸]àˆsØØ ×$Ò$ X¨tÑ4Ô4Ø×"Ò" 6¨4Ñ0Ô0Ø"×&Ò& z°4Ñ8Ô8Ø ×$Ò$ X¨tÑ4Ô4Ø(5×(9Ò(9Ø-¨tñ)ô )ð '@Ø/Ø'ð
ñ 
ô 
ð 	
rA   c                 ó˜   — |j         | j        vrt          | j        › d|j        › d�¦  «        ‚ | j        |j                  ||||¦  «        S )NzI only supports a specific set of operations, can't perform requested op (rD   )Ú_overloadpacketr!   ÚNotImplementedErrorrH   )r<   ÚfuncÚtypesÚargsÚkwargss        r?   Ú__torch_dispatch__z-SparseSemiStructuredTensor.__torch_dispatch__Ñ   sp   € àÔ sÔ':Ð:Ð:Ý%Ø”<ð @ð @Ø/3¬}ð@ð @ð @ñô ð ð 9ˆsÔ" 4Ô#7Ô8¸¸uÀdÈFÑSÔSÐSrA   c                 ó¢  — t          | dd¦  «        �€ºt          j        j        j        t
          t          j        j        j        t          t          j        j        j        t          t          j        j        j
        t          t          j        j        j        t          t          j        j        j        t          t          j        j        j        t           t          j        j        j        t$          t          j        j        j        t$          t          j        j        j        t*          t          j        j        j        t.          t          j        j        j        t2          t          j        j        j        t6          t          j        j        j        t:          t          j        j        j        t>          i| _         |�| j          !                    |¦  «         dS dS dS )zT
        Loads the op overload sparse dispatch table for the current class.
        r!   N)"rN   r6   ÚopsÚatenÚvaluesr   Úindicesr   Úis_same_sizer   Údetach_Údetachr   Útr   Úviewr   Úmmr   ÚmatmulÚaddmmr	   Úlinearr   Ú_to_copyr   Ú
_scaled_mmr   Úcloner
   Útor   r!   Úupdate)r<   Úcustom_dispatch_tables     r?   r5   z/SparseSemiStructuredTensor._load_dispatch_tableÚ   s  € õ
 �3Ð)¨4Ñ0Ô0Ñ8å”	”Ô%Õ'9Ý”	”Ô&Õ(;Ý”	”Ô+Õ-@Ý”	”Ô&Õ(;Ý”	”Ô%Õ'9Ý”	”Ô ¥-Ý”	”Ô#Õ%5Ý”	”Ô!¥>Ý”	”Ô%¥~Ý”	”Ô$Õ&7Ý”	”Ô%Õ'9Ý”	”Ô'Õ)<Ý”	”Ô)Õ+@Ý”	”Ô$Õ&7Ý”	”Ô!¥>ð#ˆCÔð" %Ð0ØÔ#×*Ò*Ð+@ÑAÔAÐAÐAÐAð' 9Ð8ð$ 1Ð0rA   Úoriginal_tensorc           	      ó.  — |j         st          d|j        › d�¦  «        ‚|                     ¦   «         dk    r%t          d|                     ¦   «         › d�¦  «        ‚|                     ¦   «         st          d¦  «        ‚|j        | j        vrt          d|j        › d| › d	�¦  «        ‚|j        \  }}| j        |j                 j        }| j        |j                 j	        }||k     s||z  s||k     s||z  rt          d
|j        › d|› d|› d�¦  «        ‚dS )z_
        Assert that the given tensor is valid for semi-structured sparse compression.
        zError original_tensor.device= z= is not supported! Only CUDA tensors are currently supported.r,   zError original_tensor.dim = z; is not supported! Only 2d tensors are currently supported.zXError original_tensor is not contiguous!Only contiguous tensors are currently supported.zError original_tensor.dtype z is not a supported dtype for ú!zError original_tensor.shape zS is not supported! Both dimensions must be larger or equal than and a multiple of (z, rD   N)
Úis_cudaÚRuntimeErrorr/   ÚdimÚis_contiguousr0   r   r)   Úsparse_min_rowsÚsparse_min_cols)r<   ry   ÚmÚnÚmin_rowsÚmin_colss         r?   Ú _validate_device_dim_dtype_shapez;SparseSemiStructuredTensor._validate_device_dim_dtype_shapeô   s�  € ð Ô&ð 	Ýð=°Ô1Gð =ð =ð =ñô ð ð ×ÒÑ Ô  AÒ%Ð%Ýð;¨×/BÒ/BÑ/DÔ/Dð ;ð ;ð ;ñô ð ð ×,Ò,Ñ.Ô.ð 	ÝðCñô ð ð Ô ¨Ô(DÐDÐDÝØj¨Ô/DÐjÐjÐdgÐjÐjÐjñô ð ð
 Ô$‰ˆˆ1ØÔ/°Ô0EÔFÔVˆØÔ/°Ô0EÔFÔVˆØˆxŠ<ˆ<˜1˜x™<ˆ<¨1¨xª<¨<¸1¸x¹<¨<åðk¨Ô/Dð kð kØS[ðkð kØ_gðkð kð kñô ð ð ,8¨<rA   c                 ó„   — | j         d         }t          j        | t          j        || j        | j        ¬¦  «        ¦  «        S )Néÿÿÿÿ©r0   r/   )r)   r6   ro   Úeyer0   r/   )rI   Úcols     r?   Úto_densez#SparseSemiStructuredTensor.to_dense  s4   € ØŒj˜ŒnˆÝŒx˜�eœi¨°4´:ÀdÄkÐRÑRÔRÑSÔSÐSrA   Úalg_idc                 ó   — t           ‚rM   ©r_   ©r<   ry   r�   s      r?   Ú
from_densez%SparseSemiStructuredTensor.from_dense#  s
   € õ "Ð!rA   )ÚbiasÚBr’   c                ó   — t           ‚rM   r�   )rI   r“   r’   rc   s       r?   Ú_mmzSparseSemiStructuredTensor._mm+  s
   € õ "Ð!rA   )Fr   FrM   )rB   N)&rH   Ú
__module__Ú__qualname__Ú__doc__r   ÚintÚ__annotations__Údictr6   r0   r   r   Úboolr   r   Ústrr   r:   rS   ÚstaticmethodÚSizer@   rJ   ÚtuplerQ   rV   Úclassmethodr\   Ú_CÚ_disabled_torch_function_implÚ__torch_function__r   rd   r5   r†   rŒ   r‘   r•   © rA   r?   r   r   *   s®  € € € € € € ðð ð" €O�SÐÐÑØ" 5¤;Ð0NÐ#NÔOÐOÐOÑOØ €N�DÐ Ð Ñ Ø!€O�TÐ!Ð!Ñ!Ø%*Ð˜dÐ*Ð*Ñ*à€L€L�LØ˜( HÐ,Ô-Ð-Ð-Ñ-àŒL˜4ÑÐÐÑØ
Œ,˜Ñ
ÐÐÑØŒl˜TÑ!Ð!Ð!Ñ!ØŒL˜4ÑÐÐÑØ!&¤°Ñ!4Ð4Ð4Ñ4Ø#Ð#Ð#Ñ#ØÐÐÑàWÐWÐW€Iàð +0Ø!"Ø#ðRð RàŒzðRð ”˜tÑ#ðRð Œl˜TÑ!ð	Rð
 ”, Ñ%ðRð ”˜tÑ#ðRð &+¤\°DÑ%8ðRð $(ðRð ðRð ðRð Rð Rñ „\ðRðh@˜#ð @ð @ð @ð @ð
*à	ˆt�CŒy˜% ¤
¨D°#°tÐ ;Ô<Ð<Ô	=ð*ð *ð *ð *ð ð
ð ˜5œ: t¨S°$Ð6Ô7ð
ð 
Œð
ð 
ð 
ñ „[ð
ð. œÔ?ÐàðT¸cð Tð Tð Tñ „[ðTð ðBð Bð Bð Bñ „[ðBð2 ð(¸u¼|ð (ÐPTð (ð (ð (ñ „[ð(ðTTð Tð Tð ð &ð"ð "àœð"ð ð"ð 
&ð	"ð "ð "ñ „[ð"ð %)ð	"ð "ð "àŒ<ð"ð Œl˜TÑ!ð	"ð 
Œð"ð "ð "ð "ð "ð "rA   r   Fry   Ú
transposedr�   rB   c                 óÈ   — |rt          j        dt          d¬¦  «         t          j        rt
          j        j        nt
          j        j        }| 	                    | |¬¦  «        S )a²	  
    This function converts a dense tensor into a sparse semi-structured tensor.
    It will return a SparseSemiStructuredTensor, a subclass of torch.Tensor.

    This function will check to ensure the dense tensor has the right dtype, size, dims, and device.
    We currently only support semi-structured sparse tensors for 2d CUDA tensors.
    Additionally, your tensor must be a positive multiple of the minimum sparse block size, given in
    `_DTYPE_TO_SHAPE_CONSTRAINTS` for each dtype (float32, float16, bfloat16, int8).

    Args:
        original_tensor (Tensor): the dense tensor to convert
        transposed (bool, optional): deprecated arg to be removed in another release. Do not use.
        alg_id (int, optional): the algorithm id to use for cuSPARSELt matmul. Defaults to 0.
            Can be obtained via ``torch._cslt_sparse_mm_search``.
    Returns:
        SparseSemiStructuredTensor: A sparse semi-structured tensor created from the given original_tensor
    Raises:
        None
    Example:
        >>> # xdoctest: +REQUIRES(env:TORCH_DOCTEST_CUDA)
        >>> A = torch.Tensor([0, 0, 1, 1]).tile((128, 32)).half().cuda()
        tensor([[0., 0., 1.,  ..., 0., 1., 1.],
                [0., 0., 1.,  ..., 0., 1., 1.],
                [0., 0., 1.,  ..., 0., 1., 1.],
                ...,
                [0., 0., 1.,  ..., 0., 1., 1.],
                [0., 0., 1.,  ..., 0., 1., 1.],
                [0., 0., 1.,  ..., 0., 1., 1.]], device='cuda:0', dtype=torch.float16)
        >>> A_sparse = to_sparse_semi_structured(A)
        SparseSemiStructuredTensor(shape=torch.Size([128, 128]))
        >>> A_sparse.values()
        tensor([[1., 1., 1.,  ..., 1., 1., 1.],
                [1., 1., 1.,  ..., 1., 1., 1.],
                [1., 1., 1.,  ..., 1., 1., 1.],
                ...,
                [1., 1., 1.,  ..., 1., 1., 1.],
                [1., 1., 1.,  ..., 1., 1., 1.],
                [1., 1., 1.,  ..., 1., 1., 1.]], device='cuda:0', dtype=torch.float16),
        >>> A_sparse.indices()
        tensor([[-4370, -4370, -4370,  ..., -4370, -4370, -4370],
                [-4370, -4370, -4370,  ..., -4370, -4370, -4370],
                [-4370, -4370, -4370,  ..., -4370, -4370, -4370],
                ...,
                [-4370, -4370, -4370,  ..., -4370, -4370, -4370],
                [-4370, -4370, -4370,  ..., -4370, -4370, -4370],
                [-4370, -4370, -4370,  ..., -4370, -4370, -4370]], device='cuda:0', dtype=torch.int16))
    z­Setting transpose from `to_sparse_semi_structured` is deprecated and will be removed in a future release. `SparseSemiStructuredTensor` only support contiguous input tensors.r,   r-   )r�   )
r2   r3   ÚFutureWarningr   r   r6   Úsparser   r   r‘   )ry   r¦   r�   ÚSPARSE_SUBCLASSs       r?   r   r   5  ss   € ðh ð 
ÝŒðRõ Øð	
ñ 	
ô 	
ð 	
õ &Ô4ð	?�ŒÔ6Ð6åŒ\Ô>ð ð ×%Ò% o¸fÐ%ÑEÔEÐErA   c                   óp  ‡ — e Zd ZdZdZej         edddd¦  «        ej         edddd¦  «        ej	         edddd¦  «        ej
         edddd¦  «        iZeej        fd	ej        d
edd fd„¦   «         Zˆ fd„Ze	 dd	ej        ddfd„¦   «         Zdddœdej        dej        dz  dedej        fd„Zˆ xZS )r   a¤  
    This class implements semi-structured sparsity for the CUTLASS backend.


    In this implementation, the specified elements and metadata are stored separately,
    in packed and meta respectively.

    When _FORCE_CUTLASS is set, or when cuSPARSELt is not available, this subclass calls into _sparse_semi_structured_(mm|addmm) and
    sparse_semi_structured_from_dense for conversion to the compressed format.
    Úcutlassé   é€   é    é@   é   é   ry   r�   rB   c           	      óˆ   — |                       |¦  «         t          |¦  «        \  }} | |j        ||d d d |j        ¬¦  «        S )N©r"   r#   r$   r%   r&   r*   )r†   r   r)   r*   )r<   ry   r�   Úsparse_tensor_cutlassÚmeta_tensor_cutlasss        r?   r‘   z,SparseSemiStructuredTensorCUTLASS.from_dense�  sf   € ð 	×,Ò,¨_Ñ=Ô=Ð=õ 6°oÑFÔFñ	
Ø!Øð ˆsØÔ!Ø(Ø$ØØØ(,Ø)Ô7ð
ñ 
ô 
ð 	
rA   c                 óÒ   •— | j         �| j        €t          d¦  «        ‚| j         j        dk    rt	          | j        | j         ¦  «        nt          ¦   «                              ¦   «         S )Nz meta and packed must not be Noner,   )r#   r"   rF   Úndimr   ÚsuperrŒ   )rI   rG   s    €r?   rŒ   z*SparseSemiStructuredTensorCUTLASS.to_dense¦  sj   ø€ ØŒ9Ð ¤Ð 3Ý Ð!CÑDÔDÐDð ŒyŒ~ Ò"Ð"õ	 4Ø”Ø”	ñô ð õ
 ‘”×!Ò!Ñ#Ô#ð	
rA   Ú r   c           	      ój   — t          j        ||d¬¦  «        \  }}}}} | |j        |||||d¬¦  «        S )aü	  
        This function takes in a unpruned dense tensor and runs a (branchless) static sort across a 4x4 tile.

        It greedily picks the largest values in the tile, upholding the 2:4 sparsity constraint across both rows and columns.
        The algorithm used to prune the matrix is implemented in `_sparse_semi_structured_tile`.

        Then it creates the packed and meta tensors for the compressed sparse representation of the pruned dense tensor.
        It also calculates the packed_t and meta_t tensors for the compressed sparse representation of the transposed
        pruned dense tensor.
        Since we cannot transpose the compressed representations, we store both for the fw/bw pass respectively.

        Finally, this function also computes a compressed swizzled bitmask that encodes the sparsity pattern.
        This can be used in the backward pass to mask the gradients.

        ::

            [9 1 7 4]                       [9 0 7 0]
            [1 2 3 0]                       [0 2 0 0]
            [8 3 5 4] -> prune 4x4 tile  -> [8 0 0 4] -> pack to CUTLASS semi-structured -> packed
            [1 2 6 2]                       [0 0 6 2]                                    -> metadata

                                                      -> pack to transposed CUTLASS      -> packed_t
                                                         semi-structured representation  -> metadata_t

                                                      -> compute swizzled bitmask        -> compressed_swizzled_bitmask

        The equivalent PyTorch code to create the same five outputs from the dense tensor can be found below::

            from torch.sparse import SparseSemiStructuredTensorCUTLASS
            from torch.sparse._semi_structured_conversions import (
                _sparse_semi_structured_tile,
                _compute_compressed_swizzled_bitmask,
            )

            pruned = _sparse_semi_structured_tile(dense)
            packed_cutlass, meta_cutlass = sparse_semi_structured_from_dense_cutlass(
                pruned
            )
            packed_t_cutlass, meta_t_cutlass = (
                sparse_semi_structured_from_dense_cutlass(pruned.t().contiguous())
            )
            bitmask = _compute_compressed_swizzled_bitmask(pruned)

            SparseSemiStructuredTensorCUTLASS(
                dense.shape,
                packed_cutlass,
                meta_cutlass,
                packed_t_cutlass,
                meta_t_cutlass,
                bitmask,
            )
        T©Ú	algorithmÚuse_cutlassFr´   )r6   Ú_sparse_semi_structured_tiler)   ©r<   ry   r½   r"   r#   r$   r%   r&   s           r?   Úprune_dense_static_sortz9SparseSemiStructuredTensorCUTLASS.prune_dense_static_sort²  sh   € õ~ Ô.Ø y¸dð
ñ 
ô 
ñ	
ØØØØØ'ð ˆsØÔ!ØØØØØ(CØð
ñ 
ô 
ð 	
rA   NF©r’   Úshould_transpose_denser“   r’   rÃ   c          
      óÊ  — t          |t          ¦  «        rt          d¦  «        ‚| j        j        }| j        dk    s|j        dk    rt          d|› d�¦  «        ‚| j        �| j        €t          d|› d�¦  «        ‚t          ¦   «          | j
        |j                 }t          j        j                             || j        | j        || j        d         |j        |j        |¦  «        S )NúZ`SparseSemiStructuredTensor @ SparseSemiStructuredTensor` is not supported by the hardwarer,   ú`ú)` matmul: Broadcasting is not implementedú$` matmul: operation is not supportedr   )Ú
isinstancer   r9   rG   rH   r¸   r_   r"   r#   Ú_ensure_cutlass_mm_registeredr   r0   r6   rf   Úsemi_structuredÚ
cutlass_mmr)   Údense_min_rowsÚdense_min_cols)rI   r“   r’   rÃ   rc   Úcls_nameÚconstraintss          r?   r•   z%SparseSemiStructuredTensorCUTLASS._mm   s   € õ �aÕ3Ñ4Ô4ð 	ÝØlñô ð ð ”>Ô*ˆØŒ9˜Š>ˆ>˜QœV qš[˜[Ý%ØG�HÐGÐGÐGñô ð ð Œ;Ð $¤)Ð"3Ý%ØB�HÐBÐBÐBñô ð õ *Ñ+Ô+Ð+ØÔ7¸¼Ô@ˆKÝ”9Ô,×7Ò7ØØ”Ø”	ØØ”
˜1”ØÔ*ØÔ*Ø&ñ	ô 	ð 	rA   ©rº   )rH   r–   r—   r˜   r    r6   Úint8r   Úfloat16Úbfloat16Úfloat32r   r¡   r   r   r:   r™   r‘   rŒ   rÁ   rœ   r•   Ú__classcell__)rG   s   @r?   r   r   |  s™  ø€ € € € € ð	ð 	ð €GàŒ
Ð2Ð2°2°s¸BÀÑCÔCØŒÐ5Ð5°b¸"¸aÀÑCÔCØŒÐ6Ð6°r¸2¸qÀ!ÑDÔDØŒÐ5Ð5°b¸"¸aÀÑCÔCð	 Ðð ð 1Ô@ð
ð 
àœð
ð ð
ð 
-ð	
ð 
ð 
ñ „[ð
ð*

ð 

ð 

ð 

ð 

ð à68ðK
ð K
Ø#œlðK
à	%ðK
ð K
ð K
ñ „[ðK
ðb %)Ø',ð!ð !ð !àŒ<ð!ð Œl˜TÑ!ð	!ð
 !%ð!ð 
Œð!ð !ð !ð !ð !ð !ð !ð !rA   r   c                   ó`  — e Zd ZdZdZej         edddd¦  «        ej         edddd¦  «        ej	         edddd¦  «        ej
         edddd¦  «        iZeej        fdej        dedd fd	„¦   «         Ze	 ddej        ddfd„¦   «         Zdddœdej        dej        dz  dedej        fd„ZdS )r   a‚  
    The cuSPARSELt backend expects the specified elements and the metadata to be stored in a single tensor:
    packed = [ specified elements of original tensor | metadata ]
    For an original tensor of size (m, k) we expect the first m * k // 2 elements to be the kept elements
    The rest of the tensor is metadata. Since there is only one tensor, we only use the packed and packed_t
    attributes respectively.

    cuSPARSELt also supports transposition fusion, which is necessary for performant 2:4 sparse training, as well
    as specifying alg_id, a config that affects the performance of the matmul depending on matmul sizes.
    Ú
cusparseltr¯   r­   r±   ry   r�   rB   c                 ó    — |                       |¦  «          | |j        t          j        |¦  «        d d d d t          j        ||j        ¬¦	  «	        S )NrX   )r†   r)   r6   Ú_cslt_compressr   r   r*   r�   s      r?   r‘   z/SparseSemiStructuredTensorCUSPARSELT.from_dense8  s`   € ð 	×,Ò,¨_Ñ=Ô=Ð=àˆsØ!Ô'ÝÔ'¨Ñ8Ô8ØØØØ(,Ý&@Ô&PØ$Ø)Ô7ð

ñ 

ô 

ð 
	
rA   rº   r   c           	      óî   — t          j        ||d¬¦  «        \  }}}}}|                     |j        d         d¦  «        }|                     |j        d         d¦  «        } | |j        |||||d¬¦  «        S )a~  
        This function does the same thing as described in SparseSemiStructuredCUTLASS, but uses the cuSPARSELt metadata
        layout and sparse matmul.

        The only functional difference is that cuSPARSELt stores `metadata` and `packed` together into a single tensor.

        ::

            [9 1 7 4]                       [9 0 7 0]
            [1 2 3 0]                       [0 2 0 0]
            [8 3 5 4] -> prune 4x4 tile  -> [8 0 0 4] -> pack to cuSPARSELT semi-structured -> packed
            [1 2 6 2]                       [0 0 6 2]

                                                      -> pack to transposed cuSPARSELt      -> packed_t
                                                         semi-structured representation

                                                      -> compute swizzled bitmask           -> compressed_swizzled_bitmask

        The equivalent PyTorch code to create the same three outputs from the dense tensor can be found below::

            from torch.sparse import SparseSemiStructuredTensorCUSPARSELT
            from torch.sparse._semi_structured_conversions import (
                _sparse_semi_structured_tile,
                _compute_compressed_swizzled_bitmask,
            )

            pruned = _sparse_semi_structured_tile(dense)
            packed_cusparselt = torch._cslt_compress(pruned)
            packed_t_cusparselt = torch._cslt_compress(pruned.t().contiguous())
            bitmask = _compute_compressed_swizzled_bitmask(pruned)

            SparseSemiStructuredTensorCUSPARSELT(
                dense.shape, packed_cutlass, None, packed_t_cutlass, None, bitmask
            )
        Fr¼   r   rˆ   é   r´   )r6   r¿   rn   r)   rÀ   s           r?   rÁ   z<SparseSemiStructuredTensorCUSPARSELT.prune_dense_static_sortL  s    € õZ Ô.Ø y¸eð
ñ 
ô 
ñ	
ØØØØØ'ð —’˜_Ô2°1Ô5°rÑ:Ô:ˆØ—=’= Ô!6°qÔ!9¸2Ñ>Ô>ˆð ˆsØÔ!ØØØØØ(CØð
ñ 
ô 
ð 	
rA   NFrÂ   r“   r’   rÃ   c                ó6  — t          |t          ¦  «        rt          d¦  «        ‚| j        dk    s|j        dk    rt	          d| j        j        › d�¦  «        ‚|j        | j        k    rWt	          d| j        j        › dt          | j	        ¦  «        › dt          |j	        ¦  «        › d| j        › d|j        › d	�¦  «        ‚|�g|j        | j        k    rWt	          d| j        j        › dt          | j	        ¦  «        › dt          |j	        ¦  «        › d
| j        › d|j        › d�¦  «        ‚| j        t          j        k    rOt	          d| j        j        › dt          | j	        ¦  «        › dt          |j	        ¦  «        › d| j        › d�	¦  «        ‚| j        €t	          d| j        j        › d�¦  «        ‚t          ¦   «          | j        |j                 }t          j        j                             || j        || j	        d         |j        |j        || j        |¦	  «	        S )NrÅ   r,   rÆ   rÇ   z` matmul: trying to do `A=z @ B=z`, with A.dtype=z and B.dtype=zH. This operation is only supported when A and B have the same data type.z + C`, with A.dtype=B.dtype=z and C.dtype=zK. This operation is only supported when A, B and C have the same data type.z`, with A.dtype=B.dtype=zO. mm is not supported for float8_e4m3fn, please use `torch._scaled_mm` instead.rÈ   r   )rÉ   r   r9   r¸   r_   rG   rH   r0   r    r)   r6   Úfloat8_e4m3fnr"   Ú _ensure_cusparselt_mm_registeredr   rf   rË   Úcusparselt_mmrÍ   rÎ   r(   )rI   r“   r’   rÃ   rc   rÐ   s         r?   r•   z(SparseSemiStructuredTensorCUSPARSELT._mm�  s·  € õ �aÕ3Ñ4Ô4ð 	ÝØlñô ð ð Œ9˜Š>ˆ>˜QœV qš[˜[Ý%ØV�D”NÔ+ÐVÐVÐVñô ð ð Œ7�d”jÒ Ð Ý%ðY�D”NÔ+ð Yð YÅuÈTÌZÑGXÔGXð Yð YÕ_dÐefÔelÑ_mÔ_mð Yð YØ $¤
ðYð YØ9:¼ðYð Yð Yñô ð ð
 Ð ¤
¨d¬jÒ 8Ð 8Ý%ð\�D”NÔ+ð \ð \ÅuÈTÌZÑGXÔGXð \ð \Õ_dÐefÔelÑ_mÔ_mð \ð \Ø(,¬
ð\ð \ØABÄð\ð \ð \ñô ð ð Œ:�Ô,Ò,Ð,Ý%ð`�D”NÔ+ð `ð `ÅuÈTÌZÑGXÔGXð `ð `Õ_dÐefÔelÑ_mÔ_mð `ð `Ø(,¬
ð`ð `ð `ñô ð ð
 Œ;ÐÝ%ØQ�D”NÔ+ÐQÐQÐQñô ð õ -Ñ.Ô.Ð.ØÔ7¸¼Ô@ˆKÝ”9Ô,×:Ò:ØØ”ØØ”
˜1”ØÔ*ØÔ*Ø&ØÔ&Ø&ñ
ô 
ð 
rA   rÑ   )rH   r–   r—   r˜   r    r6   rÞ   r   rÒ   rÓ   rÔ   r   r¡   r   r   r:   r™   r‘   rÁ   rœ   r•   r¥   rA   r?   r   r   $  so  € € € € € ð	ð 	ð €GàÔÐ;Ð;¸BÀÀBÈÑKÔKØŒ
Ð2Ð2°2°r¸2¸rÑBÔBØŒÐ5Ð5°b¸"¸aÀÑCÔCØŒÐ6Ð6°r¸2¸qÀ!ÑDÔDð	 Ðð ð 1Ô@ð
ð 
àœð
ð ð
ð 
0ð	
ð 
ð 
ñ „[ð
ð& à68ð>
ð >
Ø#œlð>
à	%ð>
ð >
ð >
ñ „[ð>
ðH %)Ø',ð4ð 4ð 4àŒ<ð4ð Œl˜TÑ!ð	4ð
 !%ð4ð 
Œð4ð 4ð 4ð 4ð 4ð 4rA   r   c                  óä  — t           rdS da ddlm}   | dd¬¦  «        dt          j        d	t          j        d
t          j        dt          j        dz  dt
          dt
          dt
          dt          dt          j        fd„¦   «         }|j        dt          j        d	t          j        d
t          j        dt          j        dz  dt
          dt
          dt
          dt          dt          j        fd„¦   «         }dS )zÄLazily register the cutlass_mm custom op.

    Registration is deferred to avoid importing torch.library at module load
    time, since torch.sparse is imported early during ``import torch``.
    NTr   ©Ú	custom_opzsemi_structured::cutlass_mmr¥   ©Úmutates_argsÚdenser"   r#   r’   Úout_featuresr„   r…   rÃ   rB   c                 óJ  — | j         \  }}	| |z  }
|	 |z  }|
dk    p|dk    }| }|r)t          j        j                             | d|d|
f¦  «        }|r|                     ¦   «         n|}|€t          j        |||¦  «        }nt          j        ||||¦  «        }|ra|r|n|	}|d |…                              dd|¦  «        }|r&|                     ¦   «          	                    ¦   «         n| 	                    ¦   «         S |r&|                     ¦   «          	                    ¦   «         n|S )Nr   rÜ   )
r)   r6   ÚnnÚ
functionalÚpadrm   Ú_sparse_semi_structured_mmÚ_sparse_semi_structured_addmmÚnarrowÚ
contiguous)ræ   r"   r#   r’   rç   r„   r…   rÃ   r‚   rƒ   Úto_pad_mÚto_pad_nÚneed_padÚdense_paddedÚmm_inputÚresÚout_colss                    r?   rÌ   z1_ensure_cutlass_mm_registered.<locals>.cutlass_mmÔ  s;  € ð Œ{‰ˆˆ1Ø�B˜(‘?ˆØ�B˜(‘?ˆØ˜q’=Ð1 H°¢MˆØˆØð 	VÝ œ8Ô.×2Ò2°5¸1¸hÈÈ8Ð:TÑUÔUˆLØ'=ÐO�<—>’>Ñ#Ô#Ð#À<ˆØˆ<ÝÔ2°6¸4ÀÑJÔJˆCˆCåÔ5°d¸FÀDÈ(ÑSÔSˆCØð 	XØ2Ð9�q�q¸ˆHØ�m�|�mÔ$×+Ò+¨A¨q°(Ñ;Ô;ˆCØ+AÐW�3—5’5‘7”7×%Ò%Ñ'Ô'Ð'ÀsÇ~Â~ÑGWÔGWÐWØ'=ÐFˆs�uŠu‰wŒw×!Ò!Ñ#Ô#Ð#À3ÐFrA   Útranspose_densec                 ó�   — |r| j         d         n| j         d         }|r||fn||f}	t          j        |	| j        | j        ¬¦  «        S ©Nr   rÜ   r‰   ©r)   r6   Úemptyr0   r/   )
ræ   r"   r#   r’   rç   r„   r…   r÷   rö   Úoutput_shapes
             r?   Ú_cutlass_mm_fakez7_ensure_cutlass_mm_registered.<locals>._cutlass_mm_fakeñ  sb   € ð &5ÐH�5”;˜q”>�>¸%¼+Àa¼.ˆà(7ÐUˆX�|Ð$Ð$¸lÈHÐ=Uð 	õ Œ{ØØ”+Ø”<ð
ñ 
ô 
ð 	
rA   )Ú_cutlass_mm_registeredÚtorch.libraryrã   r6   r:   r™   rœ   Úregister_fake)rã   rÌ   rý   s      r?   rÊ   rÊ   Ç  s]  € õ ð ØˆØ!Ðà'Ð'Ð'Ð'Ð'Ð'à€YÐ,¸2Ð>Ñ>Ô>ðGÝŒ|ðGå”ðGõ ŒlðGõ Œl˜TÑ!ð	Gõ
 ðGõ ðGõ ðGõ !%ðGõ 
ŒðGð Gð Gñ ?Ô>ðGð8 Ôð
ÝŒ|ð
å”ð
õ Œlð
õ Œl˜TÑ!ð	
õ
 ð
õ ð
õ ð
õ ð
õ 
Œð
ð 
ð 
ñ Ôð
ð 
ð 
rA   c                  óð  — t           rdS da ddlm}   | dd¬¦  «        	 dd	t          j        d
t          j        dt          j        dz  dt
          dt
          dt
          dt          dt
          dt          dt          j        fd„¦   «         }|j        d	t          j        d
t          j        dt          j        dz  dt
          dt
          dt
          dt          dt
          dt          dt          j        fd„¦   «         }dS )z,Lazily register the cusparselt_mm custom op.NTr   râ   zsemi_structured::cusparselt_mmr¥   rä   Fræ   r"   r’   rç   r„   r…   Úfuse_transposer�   rÃ   rB   c	                 ó¾  — | j         \  }	}
|	 |z  }|
 |z  }|dk    p|dk    }| }|r)t          j        j                             | d|d|f¦  «        }|r|                     ¦   «         n|}t          j        |||||¬¦  «        }|rZ|r|	n|
}|r)|                     dd|¦  «                             ¦   «         S |                     dd|¦  «                             ¦   «         S |S )Nr   )r’   Útranspose_resultr�   rÜ   )	r)   r6   ré   rê   rë   rm   Ú_cslt_sparse_mmrî   rï   )ræ   r"   r’   rç   r„   r…   r  r�   rÃ   r‚   rƒ   rð   rñ   rò   ró   rô   rõ   rö   s                     r?   rà   z7_ensure_cusparselt_mm_registered.<locals>.cusparselt_mm  s  € ð Œ{‰ˆˆ1Ø�B˜(‘?ˆØ�B˜(‘?ˆØ˜q’=Ð1 H°¢MˆØˆØð 	VÝ œ8Ô.×2Ò2°5¸1¸hÈÈ8Ð:TÑUÔUˆLØ'=ÐO�<—>’>Ñ#Ô#Ð#À<ˆÝÔ#ØØØØ+Øð
ñ 
ô 
ˆð ð 	?Ø2Ð9�q�q¸ˆHØð ?Ø—z’z ! Q¨Ñ1Ô1×<Ò<Ñ>Ô>Ð>à—z’z ! Q¨Ñ1Ô1×<Ò<Ñ>Ô>Ð>Øˆ
rA   c	                 ó�   — |r| j         d         n| j         d         }	|r|	|fn||	f}
t          j        |
| j        | j        ¬¦  «        S rù   rú   )ræ   r"   r’   rç   r„   r…   r  r�   rÃ   rö   rü   s              r?   Ú_cusparselt_mm_fakez=_ensure_cusparselt_mm_registered.<locals>._cusparselt_mm_fake6  sb   € ð &<ÐO�5”;˜q”>�>ÀÄÈQÄˆà(6ÐTˆX�|Ð$Ð$¸\È8Ð<Tð 	õ Œ{ØØ”+Ø”<ð
ñ 
ô 
ð 	
rA   )F)Ú_cusparselt_mm_registeredrÿ   rã   r6   r:   r™   rœ   r   )rã   rà   r  s      r?   rß   rß   
  sl  € õ !ð ØˆØ $Ðà'Ð'Ð'Ð'Ð'Ð'à€YÐ/¸bÐAÑAÔAð (-ð ð  ÝŒ|ð å”ð õ Œl˜TÑ!ð õ ð	 õ
 ð õ ð õ ð õ ð õ !%ð õ 
Œð ð  ð  ñ BÔAð ðD Ô ð
ÝŒ|ð
å”ð
õ Œl˜TÑ!ð
õ ð	
õ
 ð
õ ð
õ ð
õ ð
õ !%ð
õ 
Œð
ð 
ð 
ñ !Ô ð
ð 
ð 
rA   )'r2   Úcollectionsr   Úcollections.abcr   Útypingr   r6   Ú)torch.sparse._semi_structured_conversionsr   r   Ú!torch.sparse._semi_structured_opsr   r	   r
   r   r   r   r   r   r   r   r   r   r   Ú__all__r   r:   r   r   rœ   r™   r   r   r   rþ   rÊ   r  rß   r¥   rA   r?   ú<module>r     sn  ðà €€€Ø "Ð "Ð "Ð "Ð "Ð "Ø $Ð $Ð $Ð $Ð $Ð $Ø Ð Ð Ð Ð Ð à €€€ðð ð ð ð ð ð ð ðð ð ð ð ð ð ð ð ð ð ð ð ð ð ð ð ð ð ð ð ð ð ð ð ð ð ð ð ð ð"ð ð €ð ", Ø$ØCñ"ô "Ð ðH"ð H"ð H"ð H"ð H" ¤ñ H"ô H"ð H"ðZ Ø,Ô<ðDFð DFØ”\ðDFàðDFð ðDFð  ð	DFð DFð DFð DFðNeð eð eð eð eÐ(Bñ eô eð eðP]ð ]ð ]ð ]ð ]Ð+Eñ ]ô ]ð ]ð@ Ð ð=
ð =
ð =
ð@ "Ð ð@
ð @
ð @
ð @
ð @
rA   