§
    �ŠtjUi  ã            	       óæ  — d dl mZmZ d dlZd dlmZ d dlmZmZ ddlm	Z	m
Z
mZmZ er d dlmZ eeedf         z  ee         z  dz  ZneZee         dz  Zg d	¢Z eej        d
¦  «        Z eej        d¦  «        Z eej        d¦  «        Zddedededz  defd„Z eej        d¦  «        Z eej        d¦  «        Z eej         d¦  «        Z! eej"        d¦  «        Z# G d„ d¦  «        Z$d„ Z%dS )é    )ÚAnyÚTYPE_CHECKINGN)ÚTensor)Ú_add_docstrÚ_sparseé   )ÚSparseSemiStructuredTensorÚ$SparseSemiStructuredTensorCUSPARSELTÚ!SparseSemiStructuredTensorCUTLASSÚto_sparse_semi_structured)Ú_dtype.)ÚaddmmÚcheck_sparse_tensor_invariantsÚmmÚsumÚsoftmaxÚsolveÚlog_softmaxr	   r   r
   r   Úas_sparse_gradcheckaH  
sparse.addmm(mat, mat1, mat2, *, beta=1., alpha=1.) -> Tensor

This function does exact same thing as :func:`torch.addmm` in the forward,
except that it supports backward for sparse COO and CSR matrix :attr:`mat1`.
When :attr:`mat1` is a COO tensor it must have `sparse_dim = 2`.

Supports both CSR and COO storage formats.

.. note::
    **Gradient support:**

    - **COO @ Dense**: Backward is supported for both inputs. The gradient for the
      sparse input is returned as a sparse COO tensor.
    - **CSR @ Dense**: Backward is supported for both inputs. The gradient for the
      sparse input is returned as a sparse CSR tensor.
    - **CSC/BSR/BSC @ Dense**: Not supported.
    - **Sparse @ Sparse** (COO @ COO, CSR @ CSR): Forward works, but backward is
      not supported.

Args:
    mat (Tensor): a dense matrix to be added
    mat1 (Tensor): a sparse matrix to be multiplied
    mat2 (Tensor): a dense matrix to be multiplied
    beta (Number, optional): multiplier for :attr:`mat` (:math:`\beta`)
    alpha (Number, optional): multiplier for :math:`mat1 @ mat2` (:math:`\alpha`)
a  
    Performs a matrix multiplication of the sparse matrix :attr:`mat1`
    and the (sparse or strided) matrix :attr:`mat2`. Similar to :func:`torch.mm`, if :attr:`mat1` is a
    :math:`(n \times m)` tensor, :attr:`mat2` is a :math:`(m \times p)` tensor, out will be a
    :math:`(n \times p)` tensor.
    When :attr:`mat1` is a COO tensor it must have `sparse_dim = 2`.

    Supports both CSR and COO storage formats.

.. note::
    **Gradient support:**

    - **COO @ Dense**: Backward is supported for both inputs. The gradient for the
      sparse input is returned as a sparse COO tensor.
    - **CSR @ Dense**: Backward is supported for both inputs. The gradient for the
      sparse input is returned as a sparse CSR tensor.
    - **CSC/BSR/BSC @ Dense**: Not supported.
    - **Sparse @ Sparse** (COO @ COO, CSR @ CSR): Forward works, but backward is
      not supported.
    - **Mixed formats** (COO @ CSR, CSR @ COO): Not supported.

    This function also additionally accepts an optional :attr:`reduce` argument that allows
    specification of an optional reduction operation, mathematically performs the following operation:

.. math::

    z_{ij} = \bigoplus_{k = 0}^{K - 1} x_{ik} y_{kj}

where :math:`\bigoplus` defines the reduce operator. :attr:`reduce` is implemented only for
CSR storage format on CPU device.

Args:
    mat1 (Tensor): the first sparse matrix to be multiplied
    mat2 (Tensor): the second matrix to be multiplied, which could be sparse or dense
    reduce (str, optional): the reduction operation to apply for non-unique indices
        (:obj:`"sum"`, :obj:`"mean"`, :obj:`"amax"`, :obj:`"amin"`). Default :obj:`"sum"`.

Shape:
    The format of the output tensor of this function follows:
    - sparse x sparse -> sparse
    - sparse x dense -> dense

Example::

    >>> a = torch.tensor([[1., 0, 2], [0, 3, 0]]).to_sparse().requires_grad_()
    >>> a
    tensor(indices=tensor([[0, 0, 1],
                           [0, 2, 1]]),
           values=tensor([1., 2., 3.]),
           size=(2, 3), nnz=3, layout=torch.sparse_coo, requires_grad=True)
    >>> b = torch.tensor([[0, 1.], [2, 0], [0, 0]], requires_grad=True)
    >>> b
    tensor([[0., 1.],
            [2., 0.],
            [0., 0.]], requires_grad=True)
    >>> y = torch.sparse.mm(a, b)
    >>> y
    tensor([[0., 1.],
            [6., 0.]], grad_fn=<SparseAddmmBackward0>)
    >>> y.sum().backward()
    >>> a.grad
    tensor(indices=tensor([[0, 0, 1],
                           [0, 2, 1]]),
           values=tensor([1., 0., 2.]),
           size=(2, 3), nnz=3, layout=torch.sparse_coo)
    >>> c = a.detach().to_sparse_csr()
    >>> c
    tensor(crow_indices=tensor([0, 2, 3]),
           col_indices=tensor([0, 2, 1]),
           values=tensor([1., 2., 3.]), size=(2, 3), nnz=3,
           layout=torch.sparse_csr)
    >>> y1 = torch.sparse.mm(c, b, 'sum')
    >>> y1
    tensor([[0., 1.],
            [6., 0.]], grad_fn=<SparseMmReduceImplBackward0>)
    >>> y2 = torch.sparse.mm(c, b, 'max')
    >>> y2
    tensor([[0., 1.],
            [6., 0.]], grad_fn=<SparseMmReduceImplBackward0>)
aë  
sparse.sampled_addmm(input, mat1, mat2, *, beta=1., alpha=1., out=None) -> Tensor

Performs a matrix multiplication of the dense matrices :attr:`mat1` and :attr:`mat2` at the locations
specified by the sparsity pattern of :attr:`input`. The matrix :attr:`input` is added to the final result.

Mathematically this performs the following operation:

.. math::

    \text{out} = \alpha\ (\text{mat1} \mathbin{@} \text{mat2})*\text{spy}(\text{input}) + \beta\ \text{input}

where :math:`\text{spy}(\text{input})` is the sparsity pattern matrix of :attr:`input`, :attr:`alpha`
and :attr:`beta` are the scaling factors.
:math:`\text{spy}(\text{input})` has value 1 at the positions where :attr:`input` has non-zero values, and 0 elsewhere.

.. note::
    :attr:`input` must be a sparse CSR tensor. :attr:`mat1` and :attr:`mat2` must be dense tensors.

Args:
    input (Tensor): a sparse CSR matrix of shape `(m, n)` to be added and used to compute
        the sampled matrix multiplication
    mat1 (Tensor): a dense matrix of shape `(m, k)` to be multiplied
    mat2 (Tensor): a dense matrix of shape `(k, n)` to be multiplied

Keyword args:
    beta (Number, optional): multiplier for :attr:`input` (:math:`\beta`)
    alpha (Number, optional): multiplier for :math:`mat1 @ mat2` (:math:`\alpha`)
    out (Tensor, optional): output tensor. Ignored if `None`. Default: `None`.

Examples::

    >>> input = torch.eye(3, device='cuda').to_sparse_csr()
    >>> mat1 = torch.randn(3, 5, device='cuda')
    >>> mat2 = torch.randn(5, 3, device='cuda')
    >>> torch.sparse.sampled_addmm(input, mat1, mat2)
    tensor(crow_indices=tensor([0, 1, 2, 3]),
        col_indices=tensor([0, 1, 2]),
        values=tensor([ 0.2847, -0.7805, -0.1900]), device='cuda:0',
        size=(3, 3), nnz=3, layout=torch.sparse_csr)
    >>> torch.sparse.sampled_addmm(input, mat1, mat2).to_dense()
    tensor([[ 0.2847,  0.0000,  0.0000],
        [ 0.0000, -0.7805,  0.0000],
        [ 0.0000,  0.0000, -0.1900]], device='cuda:0')
    >>> torch.sparse.sampled_addmm(input, mat1, mat2, beta=0.5, alpha=0.5)
    tensor(crow_indices=tensor([0, 1, 2, 3]),
        col_indices=tensor([0, 1, 2]),
        values=tensor([ 0.1423, -0.3903, -0.0950]), device='cuda:0',
        size=(3, 3), nnz=3, layout=torch.sparse_csr)
ÚinputÚdimÚdtypeÚreturnc                 óº   — |€+|�t          j        | |¦  «        S t          j        | ¦  «        S |�t          j        | ||¬¦  «        S t          j        | |¬¦  «        S )a¥	  Return the sum of each row of the given sparse tensor.

    Returns the sum of each row of the sparse tensor :attr:`input` in the given
    dimensions :attr:`dim`. If :attr:`dim` is a list of dimensions,
    reduce over all of them. When sum over all ``sparse_dim``, this method
    returns a dense tensor instead of a sparse tensor.

    All summed :attr:`dim` are squeezed (see :func:`torch.squeeze`), resulting an output
    tensor having :attr:`dim` fewer dimensions than :attr:`input`.

    During backward, only gradients at ``nnz`` locations of :attr:`input`
    will propagate back. Note that the gradients of :attr:`input` is coalesced.

    Args:
        input (Tensor): the input sparse tensor
        dim (int or tuple of ints): a dimension or a list of dimensions to reduce. Default: reduce
            over all dims.
        dtype (:class:`torch.dtype`, optional): the desired data type of returned Tensor.
            Default: dtype of :attr:`input`.

    Example::

        >>> nnz = 3
        >>> dims = [5, 5, 2, 3]
        >>> I = torch.cat([torch.randint(0, dims[0], size=(nnz,)),
                           torch.randint(0, dims[1], size=(nnz,))], 0).reshape(2, nnz)
        >>> V = torch.randn(nnz, dims[2], dims[3])
        >>> size = torch.Size(dims)
        >>> # xdoctest: +IGNORE_WANT("non-deterministic")
        >>> S = torch.sparse_coo_tensor(I, V, size)
        >>> S
        tensor(indices=tensor([[2, 0, 3],
                               [2, 4, 1]]),
               values=tensor([[[-0.6438, -1.6467,  1.4004],
                               [ 0.3411,  0.0918, -0.2312]],

                              [[ 0.5348,  0.0634, -2.0494],
                               [-0.7125, -1.0646,  2.1844]],

                              [[ 0.1276,  0.1874, -0.6334],
                               [-1.9682, -0.5340,  0.7483]]]),
               size=(5, 5, 2, 3), nnz=3, layout=torch.sparse_coo)

        # when sum over only part of sparse_dims, return a sparse tensor
        >>> torch.sparse.sum(S, [1, 3])
        tensor(indices=tensor([[0, 2, 3]]),
               values=tensor([[-1.4512,  0.4073],
                              [-0.8901,  0.2017],
                              [-0.3183, -1.7539]]),
               size=(5, 2), nnz=3, layout=torch.sparse_coo)

        # when sum over all sparse dim, return a dense tensor
        # with summed dims squeezed
        >>> torch.sparse.sum(S, [0, 1, 3])
        tensor([-2.6596, -1.1450])
    N)r   )ÚtorchÚ_sparse_sum)r   r   r   s      úS/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/torch/sparse/__init__.pyr   r   Ù   sf   € ðr €}Øˆ?ÝÔ$ U¨CÑ0Ô0Ð0åÔ$ UÑ+Ô+Ð+àˆ?ÝÔ$ U¨C°uÐ=Ñ=Ô=Ð=åÔ$ U°%Ð8Ñ8Ô8Ð8ó    a•  
sparse.softmax(input, dim, *, dtype=None) -> Tensor

Applies a softmax function.

Softmax is defined as:

:math:`\text{Softmax}(x_{i}) = \frac{exp(x_i)}{\sum_j exp(x_j)}`

where :math:`i, j` run over sparse tensor indices and unspecified
entries are ignores. This is equivalent to defining unspecified
entries as negative infinity so that :math:`exp(x_k) = 0` when the
entry with index :math:`k` has not specified.

It is applied to all slices along `dim`, and will re-scale them so
that the elements lie in the range `[0, 1]` and sum to 1.

Args:
    input (Tensor): input
    dim (int): A dimension along which softmax will be computed.
    dtype (:class:`torch.dtype`, optional): the desired data type
        of returned tensor.  If specified, the input tensor is
        casted to :attr:`dtype` before the operation is
        performed. This is useful for preventing data type
        overflows. Default: None
a¡  
sparse.spsolve(input, other, *, left=True) -> Tensor

Computes the solution of a square system of linear equations with
a unique solution. Its purpose is similar to :func:`torch.linalg.solve`,
except that the system is defined by a sparse CSR matrix with layout
`sparse_csr`.

Args:
    input (Tensor): a sparse CSR matrix of shape `(n, n)` representing the
        coefficients of the linear system.
    other (Tensor): a dense matrix of shape `(n, )` representing the right-hand
        side of the linear system.
    left (bool, optional): whether to solve the system for `input @ out = other`
        (default) or `out @ input = other`. Only `left=True` is supported.
a  
sparse.log_softmax(input, dim, *, dtype=None) -> Tensor

Applies a softmax function followed by logarithm.

See :class:`~torch.sparse.softmax` for more details.

Args:
    input (Tensor): input
    dim (int): A dimension along which softmax will be computed.
    dtype (:class:`torch.dtype`, optional): the desired data type
        of returned tensor.  If specified, the input tensor is
        casted to :attr:`dtype` before the operation is
        performed. This is useful for preventing data type
        overflows. Default: None
a(  
sparse.spdiags(diagonals, offsets, shape, layout=None) -> Tensor

Creates a sparse 2D tensor by placing the values from rows of
:attr:`diagonals` along specified diagonals of the output

The :attr:`offsets` tensor controls which diagonals are set.

- If :attr:`offsets[i]` = 0, it is the main diagonal
- If :attr:`offsets[i]` < 0, it is below the main diagonal
- If :attr:`offsets[i]` > 0, it is above the main diagonal

The number of rows in :attr:`diagonals` must match the length of :attr:`offsets`,
and an offset may not be repeated.

Args:
    diagonals (Tensor): Matrix storing diagonals row-wise
    offsets (Tensor): The diagonals to be set, stored as a vector
    shape (2-tuple of ints): The desired shape of the result
Keyword args:
    layout (:class:`torch.layout`, optional): The desired layout of the
        returned tensor. ``torch.sparse_coo``, ``torch.sparse_csc`` and ``torch.sparse_csr``
        are supported. Default: ``torch.sparse_coo``

Examples:

Set the main and first two lower diagonals of a matrix::

    >>> diags = torch.arange(9).reshape(3, 3)
    >>> diags
    tensor([[0, 1, 2],
            [3, 4, 5],
            [6, 7, 8]])
    >>> s = torch.sparse.spdiags(diags, torch.tensor([0, -1, -2]), (3, 3))
    >>> s
    tensor(indices=tensor([[0, 1, 2, 1, 2, 2],
                           [0, 1, 2, 0, 1, 0]]),
           values=tensor([0, 1, 2, 3, 4, 6]),
           size=(3, 3), nnz=6, layout=torch.sparse_coo)
    >>> s.to_dense()
    tensor([[0, 0, 0],
            [3, 1, 0],
            [6, 4, 2]])


Change the output layout::

    >>> diags = torch.arange(9).reshape(3, 3)
    >>> diags
    tensor([[0, 1, 2],[3, 4, 5], [6, 7, 8])
    >>> s = torch.sparse.spdiags(diags, torch.tensor([0, -1, -2]), (3, 3), layout=torch.sparse_csr)
    >>> s
    tensor(crow_indices=tensor([0, 1, 3, 6]),
           col_indices=tensor([0, 0, 1, 0, 1, 2]),
           values=tensor([0, 3, 1, 6, 4, 2]), size=(3, 3), nnz=6,
           layout=torch.sparse_csr)
    >>> s.to_dense()
    tensor([[0, 0, 0],
            [3, 1, 0],
            [6, 4, 2]])

Set partial diagonals of a large output::

    >>> diags = torch.tensor([[1, 2], [3, 4]])
    >>> offsets = torch.tensor([0, -1])
    >>> torch.sparse.spdiags(diags, offsets, (5, 5)).to_dense()
    tensor([[1, 0, 0, 0, 0],
            [3, 2, 0, 0, 0],
            [0, 4, 0, 0, 0],
            [0, 0, 0, 0, 0],
            [0, 0, 0, 0, 0]])

.. note::

    When setting the values along a given diagonal the index into the diagonal
    and the index into the row of :attr:`diagonals` is taken as the
    column index in the output. This has the effect that when setting a diagonal
    with a positive offset `k` the first value along that diagonal will be
    the value in position `k` of the row of :attr:`diagonals`

Specifying a positive offset::

    >>> diags = torch.tensor([[1, 2, 3], [1, 2, 3], [1, 2, 3]])
    >>> torch.sparse.spdiags(diags, torch.tensor([0, 1, 2]), (5, 5)).to_dense()
    tensor([[1, 2, 3, 0, 0],
            [0, 2, 3, 0, 0],
            [0, 0, 3, 0, 0],
            [0, 0, 0, 0, 0],
            [0, 0, 0, 0, 0]])
c                   ón   — e Zd ZdZed„ ¦   «         Zed„ ¦   «         Zed„ ¦   «         Zdd„Zd„ Z	d„ Z
d	„ Zd
S )r   aÅ  A tool to control checking sparse tensor invariants.

    The following options exists to manage sparsr tensor invariants
    checking in sparse tensor construction:

    1. Using a context manager:

       .. code:: python

           with torch.sparse.check_sparse_tensor_invariants():
               run_my_model()

    2. Using a procedural approach:

       .. code:: python

           prev_checks_enabled = torch.sparse.check_sparse_tensor_invariants.is_enabled()
           torch.sparse.check_sparse_tensor_invariants.enable()

           run_my_model()

           if not prev_checks_enabled:
               torch.sparse.check_sparse_tensor_invariants.disable()

    3. Using function decoration:

       .. code:: python

           @torch.sparse.check_sparse_tensor_invariants()
           def run_my_model():
               ...

           run_my_model()

    4. Using ``check_invariants`` keyword argument in sparse tensor constructor call.
       For example:

       >>> torch.sparse_csr_tensor([0, 1, 3], [0, 1], [1, 2], check_invariants=True)
       Traceback (most recent call last):
         File "<stdin>", line 1, in <module>
       RuntimeError: `crow_indices[..., -1] == nnz` is not satisfied.
    c                  ó>   — t           j                             ¦   «         S )a;  Return True if the sparse tensor invariants checking is enabled.

        .. note::

            Use :func:`torch.sparse.check_sparse_tensor_invariants.enable` or
            :func:`torch.sparse.check_sparse_tensor_invariants.disable` to
            manage the state of the sparse tensor invariants checks.
        )r   Ú_CÚ_check_sparse_tensor_invariants© r   r   Ú
is_enabledz)check_sparse_tensor_invariants.is_enabledñ  s   € õ Œx×7Ò7Ñ9Ô9Ð9r   c                  óD   — t           j                             d¦  «         dS )ax  Enable sparse tensor invariants checking in sparse tensor constructors.

        .. note::

            By default, the sparse tensor invariants checks are disabled. Use
            :func:`torch.sparse.check_sparse_tensor_invariants.is_enabled` to
            retrieve the current state of sparse tensor invariants checking.

        .. note::

            The sparse tensor invariants check flag is effective to all sparse
            tensor constructors, both in Python and ATen.

        The flag can be locally overridden by the ``check_invariants``
        optional argument of the sparse tensor constructor functions.
        TN©r   r!   Ú#_set_check_sparse_tensor_invariantsr#   r   r   Úenablez%check_sparse_tensor_invariants.enableý  s    € õ$ 	Œ×4Ò4°TÑ:Ô:Ð:Ð:Ð:r   c                  óD   — t           j                             d¦  «         dS )z¯Disable sparse tensor invariants checking in sparse tensor constructors.

        See :func:`torch.sparse.check_sparse_tensor_invariants.enable` for more information.
        FNr&   r#   r   r   Údisablez&check_sparse_tensor_invariants.disable  s    € õ 	Œ×4Ò4°UÑ;Ô;Ð;Ð;Ð;r   Tc                 ó"   — || _         d | _        d S ©N)ÚstateÚsaved_state)Úselfr(   s     r   Ú__init__z'check_sparse_tensor_invariants.__init__  s   € ØˆŒ
Ø(,ˆÔÐÐr   c                 ó¬   — | j         �t          d¦  «        ‚|                      ¦   «         | _         t          j                             | j        ¦  «         d S )NzqThis context manager instance is already activated. Use a different context manager instance for context nesting.)r.   ÚRuntimeErrorr$   r   r!   r'   r-   )r/   s    r   Ú	__enter__z(check_sparse_tensor_invariants.__enter__  sU   € ØÔÐ'ÝðQñô ð ð  Ÿ?š?Ñ,Ô,ˆÔÝŒ×4Ò4°T´ZÑ@Ô@Ð@Ð@Ð@r   c                 óˆ   — | j         €t          d¦  «        ‚t          j                             | j         ¦  «         d | _         d S )Nz&saved_state should not be None on exit)r.   ÚAssertionErrorr   r!   r'   )r/   ÚtypeÚvalueÚ	tracebacks       r   Ú__exit__z'check_sparse_tensor_invariants.__exit__'  sA   € ØÔÐ#Ý Ð!IÑJÔJÐJÝŒ×4Ò4°TÔ5EÑFÔFÐFØˆÔÐÐr   c                 ó   ‡ ‡— ˆˆ fd„}|S )Nc                  ó‚   •—  t          ‰¦  «        ‰j        ¦  «        5   ‰| i |¤Žcd d d ¦  «         S # 1 swxY w Y   d S r,   )r6   r-   )ÚargsÚkwargsÚmthr/   s     €€r   Útest_mthz9check_sparse_tensor_invariants.__call__.<locals>.test_mth/  s‘   ø€ Ø•�d‘”˜DœJÑ'Ô'ð ,ð ,Ø�s˜DÐ+ FÐ+Ð+ð,ð ,ð ,ð ,ñ ,ô ,ð ,ð ,ð ,ð ,ð ,ð ,øøøð ,ð ,ð ,ð ,ð ,ð ,s   Ÿ4´8»8r#   )r/   r>   r?   s   `` r   Ú__call__z'check_sparse_tensor_invariants.__call__.  s)   øø€ ð	,ð 	,ð 	,ð 	,ð 	,ð 	,ð ˆr   N)T)Ú__name__Ú
__module__Ú__qualname__Ú__doc__Ústaticmethodr$   r(   r*   r0   r3   r9   r@   r#   r   r   r   r   Å  s´   € € € € € ð)ð )ðV ð	:ð 	:ñ „\ð	:ð ð;ð ;ñ „\ð;ð& ð<ð <ñ „\ð<ð-ð -ð -ð -ðAð Að Að ð  ð  ðð ð ð ð r   r   c                 ó   ‡ — ˆ fd„}|S )až  Decorate function, to extend gradcheck for sparse tensors.

    Decorator for torch.autograd.gradcheck or its functools.partial
    variants that extends the gradcheck function with support to input
    functions that operate on or/and return sparse tensors.

    The specified gradcheck function itself is guaranteed to operate
    on strided tensors only.

    For example:

    >>> gradcheck = torch.sparse.as_sparse_gradcheck(torch.autograd.gradcheck)
    >>> x = (
    ...     torch.tensor([[0, 1], [2, 3]], dtype=torch.float64)
    ...     .to_sparse_coo()
    ...     .requires_grad_(True)
    ... )
    >>> gradcheck(lambda x: x.to_sparse_csr(), x)
    True
    c                 ó–  •‡ ‡‡‡‡	‡
‡— |                      dd¦  «        Št          j        t          j        t          j        t          j        t          j        hŠt          j        t          j        t          j        t          j        hŠ
t          j        t          j        hŠ	dŠˆˆˆ	ˆfd„}ˆˆ
fd„Šˆ ˆˆˆfd„}| ||¦  «        f} ‰|i |¤ŽS )z©
        Create gradcheck with support for sparse tensors.

        Same as :func:`torch.autograd.gradcheck` but with sparse tensors inputs and outputs support.
        ÚmaskedFÚ__STRIDED_REPRESENTATION__c                 óx  •— t          | t          t          f¦  «        s| f} g }| D �]†}t          |t          j        ¦  «        �rS|j        �rK|j        ‰v �rA|j        |j        dœ}‰	sâ|j        | 	                    ¦   «         z
  | 
                    ¦   «         z
  }|j        ‰
v r'|                     ¦   «         j        |dz   |dz   …         nd}t          j        |j        |j        t          j        ¬¦  «                             |j        || 	                    ¦   «         ¬¦  «        }|                     ¦   «                              |¦  «        }|j        t          j        u rP|                     |                     ¦   «         |                     ¦   «         ¬¦  «         |                     ¦   «         }n¾|j        t          j        t          j        hv rP|                     |                     ¦   «         |                     ¦   «         ¬¦  «         |                     ¦   «         }nO|                     |                     ¦   «         |                     ¦   «         ¬¦  «         |                     ¦   «         }|                     ‰||                     d	¦  «        f¦  «         �Œq|                     |¦  «         �Œˆt          |¦  «        S )
ziConvert differentiable non-strided tensors to a representation containing differentiable strided tensors.)ÚlayoutÚshaper   é   N)Údevicer   )rK   Ú	blocksizeÚ	dense_dim)ÚindicesÚis_coalesced)Úcompressed_indicesÚplain_indicesT) Ú
isinstanceÚlistÚtupler   r   Úrequires_gradrK   rL   ÚndimrP   Ú
sparse_dimÚvaluesÚonesrN   ÚboolÚ	to_sparseÚto_denseÚsparse_maskÚ
sparse_cooÚupdateÚ_indicesrR   Ú_valuesÚ
sparse_csrÚ
sparse_bsrÚcrow_indicesÚcol_indicesÚccol_indicesÚrow_indicesÚextendÚrequires_grad_Úappend)r<   Únew_argsÚobjÚdÚ	batch_dimrO   Ú	full_maskr[   ÚSTRIDED_REPRESENTATIONrH   Úsparse_block_layoutsÚsparse_layoutss           €€€€r   Ú!convert_to_strided_representationzeas_sparse_gradcheck.<locals>.gradcheck_with_sparse_support.<locals>.convert_to_strided_representationc  s¥  ø€ å˜d¥T­5 MÑ2Ô2ð Ø�w�Ø"$ˆHØð 9)ñ 9)�å˜s¥E¤LÑ1Ô1ñ8)àÔ)ñ8)ð œ
 nÐ4Ñ4ð #&¤*Ø!$¤ðð �Að "ð Dà$'¤H¨s¯}ª}©¬Ñ$>ÀÇÂÑAQÔAQÑ$Q˜	ð  #œzÐ-AÐAÐAð  ŸJšJ™LœLÔ.¨y¸1©}¸yÈ1¹}Ð/LÔMÐMà!%ð "õ
 %*¤JØœI¨c¬jÅÄ
ð%ñ %ô %ç#š)Ø#&¤:Ø&/Ø&)§m¢m¡o¤oð $ñ ô ð "ð "Ÿlšl™nœn×8Ò8¸ÑCÔC˜Ø”z¥UÔ%5Ð5Ð5àŸšà$'§L¢L¡N¤Nà),×)9Ò)9Ñ);Ô);ð	 !ñ ô ð ð "%§¢¡¤˜˜Øœ­Ô(8½%Ô:JÐ'KÐKÐKàŸšà/2×/?Ò/?Ñ/AÔ/Aà*-¯/ª/Ñ*;Ô*;ð	 !ñ ô ð ð "%§¢¡¤˜˜ð Ÿšà/2×/?Ò/?Ñ/AÔ/Aà*-¯/ª/Ñ*;Ô*;ð	 !ñ ô ð ð "%§¢¡¤˜Ø—O’OØ/°°F×4IÒ4IÈ$Ñ4OÔ4OÐPñô ð ñ ð —O’O CÑ(Ô(Ð(Ñ(Ý˜‘?”?Ð"r   c                 ó(  •— g }t          | ¦  «        } | rð|                      d¦  «        }|‰k    r¾|                      d¦  «        |                      d¦  «        }}|d         t          j        u r+t          j        |d         ||d         |d         ¬¦  «        }nU|d         ‰v r2t          j        |d         |d         ||d         |d         ¬	¦  «        }nt          d
|d         › d�¦  «        ‚|                     |¦  «         | °ðt          |¦  «        S )zNRestore non-strided differentiable tensors from their strided representations.r   rK   rQ   rL   rR   )ÚsizerR   rS   rT   )rx   rK   zconversion of z! strided representation to tensor)	rV   Úpopr   ra   Úsparse_coo_tensorÚsparse_compressed_tensorÚNotImplementedErrorrm   rW   )r<   rn   Úarp   r[   rs   Úsparse_compressed_layoutss        €€r   Ú#restore_from_strided_representationzgas_sparse_gradcheck.<locals>.gradcheck_with_sparse_support.<locals>.restore_from_strided_representation¤  s6  ø€ àˆHÝ˜‘:”:ˆDØð #Ø—H’H˜Q‘K”K�ØÐ.Ò.Ð.Ø $§¢¨¡¤¨T¯XªX°a©[¬[�v�AØ˜”{¥eÔ&6Ð6Ð6Ý!Ô3Ø˜iœLØ"Ø!" 7¤Ø)*¨>Ô):ð	ñ ô ˜˜ð ˜8œÐ(AÐAÐAÝ!Ô:ØÐ2Ô3Ø˜oÔ.Ø"Ø!" 7¤Ø#$ X¤;ðñ ô ˜˜õ 2Ø[¨Q¨x¬[Ð[Ð[Ð[ñô ð ð —’ Ñ"Ô"Ð"ð/ ð #õ0 ˜‘?”?Ð"r   c                  ó
  •—  ‰| ¦  «        } ‰|i |¤Ž}t          |t          t          f¦  «        rt          |¦  «        n|f}t          ˆˆfd„|D ¦   «         ¦  «        }t          |t          t          f¦  «        r|n|d         S )Nc              3   óœ   •K  — | ]F}t          |t          j        ¦  «        r&|j        r|j        ‰v r|                     ‰¬ ¦  «        n|V — ŒGdS ))Úmasked_gradN)rU   r   r   rX   rK   r_   )Ú.0ÚorH   ru   s     €€r   ú	<genexpr>zcas_sparse_gradcheck.<locals>.gradcheck_with_sparse_support.<locals>.func_wrapper.<locals>.<genexpr>Ì  s~   øè è € ð 	$ð 	$ð õ " !¥U¤\Ñ2Ô2ðàœðð œ NÐ2Ð2ð —J’J¨6�JÑ2Ô2Ð2ð ð	$ð 	$ð 	$ð 	$ð 	$ð 	$r   r   )rU   rV   rW   )	r<   r=   Úrestored_argsÚoutputsÚstrided_outputsÚfuncrH   r   ru   s	        €€€€r   Úfunc_wrapperzPas_sparse_gradcheck.<locals>.gradcheck_with_sparse_support.<locals>.func_wrapperÂ  s¹   ø€ Ø?Ð?ÀÑEÔEˆMð �d˜MÐ4¨VÐ4Ð4ˆGõ #-¨Wµt½U°mÑ"DÔ"DÐT•�g‘”�È7È*ð õ $ð 	$ð 	$ð 	$ð 	$ð 	$ð )ð	$ñ 	$ô 	$ñ 	ô 	ˆOõ ˜g­­e }Ñ5Ô5ð(��à$ QÔ'ðr   )ry   r   ra   re   Ú
sparse_cscrf   Ú
sparse_bsc)r‰   Úinputsr=   rv   rŠ   r<   rs   rH   r   rt   r~   ru   Ú	gradchecks   `     @@@@@@€r   Úgradcheck_with_sparse_supportz:as_sparse_gradcheck.<locals>.gradcheck_with_sparse_supportL  s  øøøøøøøø€ ð —’˜H eÑ,Ô,ˆåÔÝÔÝÔÝÔÝÔð
ˆõ ÔÝÔÝÔÝÔð	%
Ð!õ !&Ô 0µ%Ô2BÐCÐØ!=Ðð?	#ð ?	#ð ?	#ð ?	#ð ?	#ð ?	#ð ?	#ð ?	#ðB	#ð 	#ð 	#ð 	#ð 	#ð 	#ð<	ð 	ð 	ð 	ð 	ð 	ð 	ð 	ð6 Ð?Ð?ÀÑGÔGÐHˆàˆy˜$Ð) &Ð)Ð)Ð)r   r#   )rŽ   r�   s   ` r   r   r   6  s*   ø€ ð,S*ð S*ð S*ð S*ð S*ðj )Ð(r   )NN)&Útypingr   r   r   r   Útorch._Cr   r   Úsemi_structuredr	   r
   r   r   Útorch.typesr   ÚDTypeÚintrW   rV   Ú	DimOrDimsÚ__all__Ú_sparse_addmmr   Ú
_sparse_mmr   Úsparse_sampled_addmmÚsampled_addmmr   Ú_sparse_softmaxr   Ú_spsolveÚspsolveÚ_sparse_log_softmaxr   Ú_spdiagsÚspdiagsr   r   r#   r   r   ú<module>r¢      s`  ðð &Ð %Ð %Ð %Ð %Ð %Ð %Ð %à €€€Ø Ð Ð Ð Ð Ð Ø )Ð )Ð )Ð )Ð )Ð )Ð )Ð )ðð ð ð ð ð ð ð ð ð ð ð ð ð "Ø+Ð+Ð+Ð+Ð+Ð+à�e˜C ˜H”oÑ%¨¨S¬	Ñ1°DÑ8€I€Ið €EØ�c”
˜TÑ!€Iðð ð €ð  	ˆØÔðñ	ô 	€ð@ €[ØÔðOñRô R€ðj �ØÔ ð1ñ4ô 4€ðnB9ð B9ˆvð B9˜Ið B9°U¸T±\ð B9ÈVð B9ð B9ð B9ð B9ðJ ˆ+ØÔðñô €ð> ˆ+ØÔðñô €ð( ˆkØÔðñô €ð* ˆ+ØÔðYñ\ô \€ð~nð nð nð nð nñ nô nð nðbk)ð k)ð k)ð k)ð k)r   