
    JThӒ                        S r SSKrSSKrSSKrSSKJr  SSKJrJrJ	r	J
r
JrJr  SSKJrJr  SSKrSSKJr  \" \R&                  5      \" S5      :  a  \" S5      eCSS	KJr  SSKJs  Jr  SS
KJr  / SQr\R:                  " SSS9r\R?                  5         \" S5      r \" S5      r!\" S5      r"\" S5      r#\r$\r%\\%/\&\'\   \$4   4   r(\\\   \$/\%4   r)\\$\\   /\%4   r*\r+\\$/\+4   r,\\+/\$4   r-\&\S4   r.\\%/\&\'\&\\4      \4   4   r/S\)S\*4S jr0SSSSS.S\1\   S\(S\)S\	\2   S\	\,   S\	\-   S\	\/   SS4S jjr3\" S \4S!9SSSS".S\1\   S\(S\)S\	\2   S\	\,   S\	\-   SS4S# jj5       r5SSSS".S\1\   S\(S\)S\	\2   S\	\,   S\	\-   SS4S$ jjr6S%\S\\   4S& jr7 SlS'\%S(\	\\%/\84      S\84S) jjr9 SlS'\%S(\	\\%/\84      S\&\'\   \4   4S* jjr:S+\\   S,\S\%4S- jr; SlS'\%S(\	\\%/\84      S\\   4S. jjr< SlS'\%S(\	\\%/\84      S\'\   4S/ jjr= SlS'\%S(\	\\%/\84      S\4S0 jjr>SS1.S\S\4   S'\%S2\%S(\	\\%/\84      S\%4
S3 jjr?SS1.S\S\4   S'\%S2\%S(\	\\%/\84      S\%4
S4 jjr@\&\1\    \1\!   4   rA\&\1\    \1\!   \1\"   4   rB\R                  S5:  a  \\1\   \&\1\   S4   \R                  4   rEO\\1\   \&\1\   S4   4   rE\\\ \!4   /\#4   rF\\\ \!\"4   /\#4   rG\\ /\#4   rH\\/\#4   rI\\ /\\/\4   4   rJ\
S6\1\    S\J\H\ \4      4S7 j5       rK\
S6\A\ \!4   S\J\F\ \!\4      4S8 j5       rK\
S6\B\ \!\"4   S\J\G\ \!\"\4      4S9 j5       rK\
S6\ES\J\I\      4S: j5       rK\
S6\\/\84   S\J\I\      4S; j5       rKS6\\E\\/\84   4   S\J\I\      4S< jrK\
 SlS\H\ \4   S'\%S(\	\\%/\84      S6\1\    S\%4
S= jj5       rL\
 SlS\F\ \!\4   S'\%S(\	\\%/\84      S6\A\ \!4   S\%4
S> jj5       rL\
 SlS\G\ \!\"\4   S'\%S(\	\\%/\84      S6\B\ \!\"4   S\%4
S? jj5       rL\
 SlS\I\   S'\%S(\	\\%/\84      S6\ES\%4
S@ jj5       rL\
 SlS\I\   S'\%S(\	\\%/\84      S6\\/\84   S\%4
SA jj5       rL SlS\I\   S'\%S(\	\\%/\84      S6\\E\\/\84   4   S\%4
SB jjrL\
 SlS\H\ \4   S'\%S(\	\\%/\84      S6\1\    S\%4
SC jj5       rM\
 SlS\F\ \!\4   S'\%S(\	\\%/\84      S6\A\ \!4   S\%4
SD jj5       rM\
 SlS\G\ \!\"\4   S'\%S(\	\\%/\84      S6\B\ \!\"4   S\%4
SE jj5       rM\
 SlS\I\   S'\%S(\	\\%/\84      S6\ES\%4
SF jj5       rM\
 SlS\I\   S'\%S(\	\\%/\84      S6\\/\84   S\%4
SG jj5       rM SlS\I\   S'\%S(\	\\%/\84      S6\\E\\/\84   4   S\%4
SH jjrM SlSI\\/\84   S'\%S(\	\\%/\84      S\84SJ jjrN SlSI\\/\84   S'\%S(\	\\%/\84      S\84SK jjrO\
 SlSI\H\ \84   S'\%S(\	\\%/\84      SL\1\    S\84
SM jj5       rP\
 SlSI\F\ \!\84   S'\%S(\	\\%/\84      SL\A\ \!4   S\84
SN jj5       rP\
 SlSI\G\ \!\"\84   S'\%S(\	\\%/\84      SL\B\ \!\"4   S\84
SO jj5       rP SlSI\I\8   S'\%S(\	\\%/\84      SL\ES\84
SP jjrP\
 SlSI\H\ \84   S'\%S(\	\\%/\84      SL\1\    S\84
SQ jj5       rQ\
 SlSI\F\ \!\84   S'\%S(\	\\%/\84      SL\A\ \!4   S\84
SR jj5       rQ\
 SlSI\G\ \!\"\84   S'\%S(\	\\%/\84      SL\B\ \!\"4   S\84
SS jj5       rQ SlSI\I\8   S'\%S(\	\\%/\84      SL\ES\84
ST jjrQ SlSU\%SV\%S(\	\\%/\84      S\'\   4SW jjrR SlS'\%S,\S(\	\\%/\84      S\	\'\      4SX jjrSSlS,\SY\	\T   S\24SZ jjrU\R                  S[\2S\4S\ j5       rW " S] S^5      rXS,\S\24S_ jrY " S` Sa\1" \5      5      rZ " Sb Sc\\ZSd9r[ SlS'\%S(\	\\%/\84      S\&\'\&\.\4      \4   4Se jjr\ SlS'\%S(\	\\%/\84      S\'\&\.\4      4Sf jjr]SS1.S\S\4   S'\%S2\%S(\	\\%/\84      S\%4
Sg jjr^Sh\.S\24Si jr_S%\Sh\.S\4Sj jr`\R                     S\lb        Sk0 srcrd\R                   H  u  rcrd\6" \c0 \dD6  M     \R                  R                  5         CcCdSSS5        g! , (       d  f       g= f)ma  
Contains utility functions for working with nested python data structures.

A *pytree* is Python nested data structure. It is a tree in the sense that
nodes are Python collections (e.g., list, tuple, dict) and the leaves are
Python values. Furthermore, a pytree should not contain reference cycles.

pytrees are useful for working with nested collections of Tensors. For example,
one can use `tree_map` to map a function over all Tensors inside some nested
collection of Tensors and `tree_leaves` to get a flat list of all Tensors
inside some nested collection. pytrees are helpful for implementing nested
collection support for PyTorch APIs.
    N)Iterable)AnyCallableOptionaloverloadTypeVarUnion)
deprecatedTypeIs)Versionz0.13.0ztorch.utils._cxx_pytree depends on optree, which is an optional dependency of PyTorch. To use it, please upgrade your optree package to >= 0.13.0)
PyTreeSpec)KeyEntry)PyTreeContextFlattenFuncUnflattenFuncDumpableContextToDumpableContextFnFromDumpableContextFnTreeSpecLeafSpeckeystrkey_getregister_pytree_nodetree_flattentree_flatten_with_pathtree_unflatten	tree_itertree_leavestree_leaves_with_pathtree_structuretree_maptree_map_with_path	tree_map_tree_map_onlytree_map_only_tree_alltree_anytree_all_onlytree_any_onlytreespec_dumpstreespec_loadstreespec_pprintTtorch	namespaceTSUR.funcreturnc                 n   ^  [         R                  " T 5      S[        S[        S[        4U 4S jj5       nU$ )Nargskwargsr6   c                  &   > T" [        U 5      0 UD6$ N)reversed)r8   r9   r5   s     O/var/www/auris/envauris/lib/python3.13/site-packages/torch/utils/_cxx_pytree.pywrapped_reverse_args.<locals>.wrappede   s    Xd^.v..    )	functoolswrapsr   )r5   r>   s   ` r=   _reverse_argsrC   d   s:    __T/s /c /c / / Nr@   )serialized_type_nameto_dumpable_contextfrom_dumpable_contextflatten_with_keys_fncls
flatten_fnunflatten_fnrD   rE   rF   rG   c          	      n    Ub  [        S5      e[        U UUUUUS9  [        R                  " U UUUUUS9  g)a  Register a container-like type as pytree node.

Args:
    cls (type): A Python type to treat as an internal pytree node.
    flatten_fn (callable): A function to be used during flattening, taking an instance of
        ``cls`` and returning a pair, with (1) an iterable for the children to be flattened
        recursively, and (2) some hashable auxiliary data to be stored in the treespec and to be
        passed to the ``unflatten_fn``.
    unflatten_fn (callable): A function taking two arguments: the auxiliary data that was
        returned by ``flatten_fn`` and stored in the treespec, and the unflattened children.
        The function should return an instance of ``cls``.
    serialized_type_name (str, optional): A keyword argument used to specify the fully
        qualified name used when serializing the tree spec.
    to_dumpable_context (callable, optional): An optional keyword argument to custom specify how
        to convert the context of the pytree to a custom json dumpable representation. This is
        used for json serialization, which is being used in :mod:`torch.export` right now.
    from_dumpable_context (callable, optional): An optional keyword argument to custom specify
        how to convert the custom json dumpable representation of the context back to the
        original context. This is used for json deserialization, which is being used in
        :mod:`torch.export` right now.

Example::

    >>> # xdoctest: +SKIP
    >>> # Registry a Python type with lambda functions
    >>> register_pytree_node(
    ...     set,
    ...     lambda s: (sorted(s), None, None),
    ...     lambda children, _: set(children),
    ... )
N-KeyPaths are not yet supported in cxx_pytree.rD   rE   rF   )NotImplementedError_private_register_pytree_nodepython_pytree)rH   rI   rJ   rD   rE   rF   rG   s          r=   r   r   l   sS    R '!"QRR!1/3 //1/3r@   z`torch.utils._cxx_pytree._register_pytree_node` is deprecated. Please use `torch.utils._cxx_pytree.register_pytree_node` instead.)categoryrM   c          	           [        U UUUUUS9  g)a3  Register a container-like type as pytree node for the C++ pytree only.

The ``namespace`` argument is used to avoid collisions that occur when different libraries
register the same Python type with different behaviors. It is recommended to add a unique prefix
to the namespace to avoid conflicts with other libraries. Namespaces can also be used to specify
the same class in different namespaces for different use cases.

.. warning::
    For safety reasons, a ``namespace`` must be specified while registering a custom type. It is
    used to isolate the behavior of flattening and unflattening a pytree node type. This is to
    prevent accidental collisions between different libraries that may register the same type.

Args:
    cls (type): A Python type to treat as an internal pytree node.
    flatten_fn (callable): A function to be used during flattening, taking an instance of
        ``cls`` and returning a pair, with (1) an iterable for the children to be flattened
        recursively, and (2) some hashable auxiliary data to be stored in the treespec and to be
        passed to the ``unflatten_fn``.
    unflatten_fn (callable): A function taking two arguments: the auxiliary data that was
        returned by ``flatten_fn`` and stored in the treespec, and the unflattened children.
        The function should return an instance of ``cls``.
    serialized_type_name (str, optional): A keyword argument used to specify the fully
        qualified name used when serializing the tree spec.
    to_dumpable_context (callable, optional): An optional keyword argument to custom specify how
        to convert the context of the pytree to a custom json dumpable representation. This is
        used for json serialization, which is being used in :mod:`torch.export` right now.
    from_dumpable_context (callable, optional): An optional keyword argument to custom specify
        how to convert the custom json dumpable representation of the context back to the
        original context. This is used for json deserialization, which is being used in
        :mod:`torch.export` right now.
rM   N)rO   rH   rI   rJ   rD   rE   rF   s         r=   _register_pytree_noderT      s    \ "1/3r@   c                |    [         R                  " U 5      (       d!  [         R                  " U U[        U5      SS9  gg)zThis is an internal function that is used to register a pytree node type
for the C++ pytree only. End-users should use :func:`register_pytree_node`
instead.
r.   r/   N)optreeis_structseq_classr   rC   rS   s         r=   rO   rO      s9     $$S))##,'		
 *r@   objc                "    [        U [        5      $ r;   )
isinstancer   )rX   s    r=   _is_pytreespec_instancer[      s    c8$$r@   treeis_leafc                 0    [         R                  " U USSS9$ )ax  Check if a pytree is a leaf.

>>> tree_is_leaf(1)
True
>>> tree_is_leaf(None)
True
>>> tree_is_leaf([1, 2, 3])
False
>>> tree_is_leaf((1, 2, 3), is_leaf=lambda x: isinstance(x, tuple))
True
>>> tree_is_leaf({'a': 1, 'b': 2, 'c': 3})
False
>>> tree_is_leaf({'a': 1, 'b': 2, 'c': None})
False

Args:
    tree (pytree): A pytree to check if it is a leaf node.
    is_leaf (callable, optional): An extra leaf predicate function that will be called at each
        flattening step. The function should have a single argument with signature
        ``is_leaf(node) -> bool``. If it returns :data:`True`, the whole subtree being treated
        as a leaf. Otherwise, the default pytree registry will be used to determine a node is a
        leaf or not. If the function is not specified, the default pytree registry will be used.

Returns:
    A boolean indicating if the pytree is a leaf node.
Tr.   r]   none_is_leafr0   )rV   tree_is_leafr\   r]   s     r=   ra   ra      s#    < 	 r@   c                 0    [         R                  " U USSS9$ )a  Flatten a pytree.

See also :func:`tree_unflatten`.

The flattening order (i.e., the order of elements in the output list) is deterministic,
corresponding to a left-to-right depth-first tree traversal.

>>> tree = {"b": (2, [3, 4]), "a": 1, "c": None, "d": 5}
>>> tree_flatten(tree)
([2, 3, 4, 1, None, 5], PyTreeSpec({'b': (*, [*, *]), 'a': *, 'c': *, 'd': *}, NoneIsLeaf, namespace='torch'))
>>> tree_flatten(1)
([1], PyTreeSpec(*, NoneIsLeaf, namespace='torch'))
>>> tree_flatten(None)
([None], PyTreeSpec(*, NoneIsLeaf, namespace='torch'))
>>> from collections import OrderedDict
>>> tree = OrderedDict([("b", (2, [3, 4])), ("a", 1), ("c", None), ("d", 5)])
>>> tree_flatten(tree)
([2, 3, 4, 1, None, 5], PyTreeSpec(OrderedDict({'b': (*, [*, *]), 'a': *, 'c': *, 'd': *}), NoneIsLeaf, namespace='torch'))

Args:
    tree (pytree): A pytree to flatten.
    is_leaf (callable, optional): An extra leaf predicate function that will be called at each
        flattening step. The function should have a single argument with signature
        ``is_leaf(node) -> bool``. If it returns :data:`True`, the whole subtree being treated
        as a leaf. Otherwise, the default pytree registry will be used to determine a node is a
        leaf or not. If the function is not specified, the default pytree registry will be used.

Returns:
    A pair ``(leaves, treespec)`` where the first element is a list of leaf values and the
    second element is a treespec representing the structure of the pytree.
Tr.   r_   )rV   r   rb   s     r=   r   r   %  s$    F 	 r@   leavestreespecc                 ~    [        U5      (       d  [        S[        U5       S35      e[        R                  " X5      $ )a0  Reconstruct a pytree from the treespec and the leaves.

The inverse of :func:`tree_flatten`.

>>> tree = {"b": (2, [3, 4]), "a": 1, "c": None, "d": 5}
>>> leaves, treespec = tree_flatten(tree)
>>> tree == tree_unflatten(leaves, treespec)
True

Args:
    leaves (iterable): The list of leaves to use for reconstruction. The list must match the
        number of leaves of the treespec.
    treespec (TreeSpec): The treespec to reconstruct.

Returns:
    The reconstructed pytree, containing the ``leaves`` placed in the structure described by
    ``treespec``.
zhtree_unflatten(leaves, treespec): Expected `treespec` to be instance of PyTreeSpec but got item of type .)r[   	TypeErrortyperV   r   )rd   re   s     r=   r   r   P  sF    & #8,,//3H~.>aA
 	
   22r@   c                 0    [         R                  " U USSS9$ )a$  Get an iterator over the leaves of a pytree.

See also :func:`tree_flatten`.

>>> tree = {"b": (2, [3, 4]), "a": 1, "c": None, "d": 5}
>>> list(tree_iter(tree))
[2, 3, 4, 1, None, 5]
>>> list(tree_iter(1))
[1]
>>> list(tree_iter(None))
[None]

Args:
    tree (pytree): A pytree to flatten.
    is_leaf (callable, optional): An extra leaf predicate function that will be called at each
        flattening step. The function should have a single argument with signature
        ``is_leaf(node) -> bool``. If it returns :data:`True`, the whole subtree being treated
        as a leaf. Otherwise, the default pytree registry will be used to determine a node is a
        leaf or not. If the function is not specified, the default pytree registry will be used.

Returns:
    An iterator over the leaf values.
Tr.   r_   )rV   r   rb   s     r=   r   r   k  s#    6 	 r@   c                 0    [         R                  " U USSS9$ )a  Get the leaves of a pytree.

See also :func:`tree_flatten`.

>>> tree = {"b": (2, [3, 4]), "a": 1, "c": None, "d": 5}
>>> tree_leaves(tree)
[2, 3, 4, 1, None, 5]
>>> tree_leaves(1)
[1]
>>> tree_leaves(None)
[None]

Args:
    tree (pytree): A pytree to flatten.
    is_leaf (callable, optional): An extra leaf predicate function that will be called at each
        flattening step. The function should have a single argument with signature
        ``is_leaf(node) -> bool``. If it returns :data:`True`, the whole subtree being treated
        as a leaf. Otherwise, the default pytree registry will be used to determine a node is a
        leaf or not. If the function is not specified, the default pytree registry will be used.

Returns:
    A list of leaf values.
Tr.   r_   )rV   r   rb   s     r=   r   r     s#    6 	 r@   c                 0    [         R                  " U USSS9$ )a  Get the treespec for a pytree.

See also :func:`tree_flatten`.

>>> tree = {"b": (2, [3, 4]), "a": 1, "c": None, "d": 5}
>>> tree_structure(tree)
PyTreeSpec({'b': (*, [*, *]), 'a': *, 'c': *, 'd': *}, NoneIsLeaf, namespace='torch')
>>> tree_structure(1)
PyTreeSpec(*, NoneIsLeaf, namespace='torch')
>>> tree_structure(None)
PyTreeSpec(*, NoneIsLeaf, namespace='torch')

Args:
    tree (pytree): A pytree to flatten.
    is_leaf (callable, optional): An extra leaf predicate function that will be called at each
        flattening step. The function should have a single argument with signature
        ``is_leaf(node) -> bool``. If it returns :data:`True`, the whole subtree being treated
        as a leaf. Otherwise, the default pytree registry will be used to determine a node is a
        leaf or not. If the function is not specified, the default pytree registry will be used.

Returns:
    A treespec object representing the structure of the pytree.
Tr.   r_   )rV   r!   rb   s     r=   r!   r!     s#    6   	 r@   r]   restsc                <    [         R                  " U U/UQ7USSS.6$ )a  Map a multi-input function over pytree args to produce a new pytree.

See also :func:`tree_map_`.

>>> tree_map(lambda x: x + 1, {"x": 7, "y": (42, 64)})
{'x': 8, 'y': (43, 65)}
>>> tree_map(lambda x: x is None, {"x": 7, "y": (42, 64), "z": None})
{'x': False, 'y': (False, False), 'z': True}

If multiple inputs are given, the structure of the tree is taken from the first input;
subsequent inputs need only have ``tree`` as a prefix:

>>> tree_map(lambda x, y: [x] + y, [5, 6], [[7, 9], [1, 2]])
[[5, 7, 9], [6, 1, 2]]

Args:
    func (callable): A function that takes ``1 + len(rests)`` arguments, to be applied at the
        corresponding leaves of the pytrees.
    tree (pytree): A pytree to be mapped over, with each leaf providing the first positional
        argument to function ``func``.
    rests (tuple of pytree): A tuple of pytrees, each of which has the same structure as
        ``tree`` or has ``tree`` as a prefix.
    is_leaf (callable, optional): An extra leaf predicate function that will be called at each
        flattening step. The function should have a single argument with signature
        ``is_leaf(node) -> bool``. If it returns :data:`True`, the whole subtree being treated
        as a leaf. Otherwise, the default pytree registry will be used to determine a node is a
        leaf or not. If the function is not specified, the default pytree registry will be used.

Returns:
    A new pytree with the same structure as ``tree`` but with the value at each leaf given by
    ``func(x, *xs)`` where ``x`` is the value at the corresponding leaf in ``tree`` and ``xs``
    is the tuple of values at corresponding nodes in ``rests``.
Tr.   r_   )rV   r"   r5   r\   r]   rn   s       r=   r"   r"     s6    N ?? 
  r@   c                <    [         R                  " U U/UQ7USSS.6$ )a  Like :func:`tree_map`, but do an inplace call on each leaf and return the original tree.

See also :func:`tree_map`.

Args:
    func (callable): A function that takes ``1 + len(rests)`` arguments, to be applied at the
        corresponding leaves of the pytrees.
    tree (pytree): A pytree to be mapped over, with each leaf providing the first positional
        argument to function ``func``.
    rests (tuple of pytree): A tuple of pytrees, each of which has the same structure as
        ``tree`` or has ``tree`` as a prefix.
    is_leaf (callable, optional): An extra leaf predicate function that will be called at each
        flattening step. The function should have a single argument with signature
        ``is_leaf(node) -> bool``. If it returns :data:`True`, the whole subtree being treated
        as a leaf. Otherwise, the default pytree registry will be used to determine a node is a
        leaf or not. If the function is not specified, the default pytree registry will be used.

Returns:
    The original ``tree`` with the value at each leaf is given by the side-effect of function
    ``func(x, *xs)`` (not the return value) where ``x`` is the value at the corresponding leaf
    in ``tree`` and ``xs`` is the tuple of values at values at corresponding nodes in ``rests``.
Tr.   r_   )rV   r$   rp   s       r=   r$   r$     s7    8  
  r@      
   type_or_types_or_predc                    g r;    ru   s    r=   map_onlyry   <      r@   c                    g r;   rw   rx   s    r=   ry   ry   A  rz   r@   c                    g r;   rw   rx   s    r=   ry   ry   F  rz   r@   c                    g r;   rw   rx   s    r=   ry   ry   L  rz   r@   c                    g r;   rw   rx   s    r=   ry   ry   Q  rz   r@   c                l  ^ ^ [        T [        [        45      (       d3  [        R                  S:  a4  [        T [
        R                  5      (       a  S[        S[        4U 4S jjmO[        T 5      (       a  T mO[        S5      eS[        [        /[        4   S[        [        /[        4   4U4S jjnU$ )aw  
Suppose you are writing a tree_map over tensors, leaving everything
else unchanged.  Ordinarily you would have to write:

    def go(t):
        if isinstance(t, Tensor):
            return ...
        else:
            return t

With this function, you only need to write:

    @map_only(Tensor)
    def go(t):
        return ...

You can also directly use 'tree_map_only'
rr   xr6   c                    > [        U T5      $ r;   rZ   )r   ru   s    r=   predmap_only.<locals>.predp  s    a!677r@   z9Argument must be a type, a tuple of types, or a callable.r5   c                 f   >^  [         R                  " T 5      S[        S[        4U U4S jj5       nU$ )Nr   r6   c                 2   > T" U 5      (       a  T" U 5      $ U $ r;   rw   )r   r5   r   s    r=   r>   *map_only.<locals>.wrapper.<locals>.wrappedy  s    AwwAwHr@   )rA   rB   r1   r   )r5   r>   r   s   ` r=   wrappermap_only.<locals>.wrapperx  s3    			q 	S 	 
	
 r@   )rZ   ri   tuplesysversion_infotypes	UnionTyper   boolcallablerh   r   r1   )ru   r   r   s   ` @r=   ry   ry   V  s    * '$77G#,eoo>>	8C 	8D 	8 	8 
'	(	($STThsCx( XseSj-A  Nr@   c                    g r;   rw   ru   r5   r\   r]   s       r=   r%   r%          r@   c                    g r;   rw   r   s       r=   r%   r%     r   r@   c                    g r;   rw   r   s       r=   r%   r%     r   r@   c                    g r;   rw   r   s       r=   r%   r%     r   r@   c                    g r;   rw   r   s       r=   r%   r%     r   r@   c                4    [        [        U 5      " U5      X#S9$ Nrm   )r"   ry   r   s       r=   r%   r%     s     H23D94QQr@   c                    g r;   rw   r   s       r=   r&   r&     r   r@   c                    g r;   rw   r   s       r=   r&   r&     r   r@   c                    g r;   rw   r   s       r=   r&   r&     r   r@   c                    g r;   rw   r   s       r=   r&   r&     r   r@   c                    g r;   rw   r   s       r=   r&   r&     r   r@   c                4    [        [        U 5      " U5      X#S9$ r   )r$   ry   r   s       r=   r&   r&     s     X34T:DRRr@   r   c                 <    [        XS9n[        [        X5      5      $ r   )r   allmapr   r\   r]   	flat_argss       r=   r'   r'         
 $0Is4#$$r@   c                 <    [        XS9n[        [        X5      5      $ r   )r   anyr   r   s       r=   r(   r(     r   r@   type_or_typesc                    g r;   rw   r   r   r\   r]   s       r=   r)   r)     r   r@   c                    g r;   rw   r   s       r=   r)   r)   #  r   r@   c                    g r;   rw   r   s       r=   r)   r)   .  r   r@   c                D   ^ ^ [        X#S9n[        UU 4S jU 5       5      $ )Nrm   c              3   Z   >#    U  H   n[        UT5      (       d  M  T" U5      v   M"     g 7fr;   r   .0r   r   r   s     r=   	<genexpr> tree_all_only.<locals>.<genexpr>A  "     J	1Z=-IwtAww	   ++)r   r   r   r   r\   r]   r   s   ``   r=   r)   r)   9        $0IJ	JJJr@   c                    g r;   rw   r   s       r=   r*   r*   D  r   r@   c                    g r;   rw   r   s       r=   r*   r*   O  r   r@   c                    g r;   rw   r   s       r=   r*   r*   Z  r   r@   c                D   ^ ^ [        X#S9n[        UU 4S jU 5       5      $ )Nrm   c              3   Z   >#    U  H   n[        UT5      (       d  M  T" U5      v   M"     g 7fr;   r   r   s     r=   r    tree_any_only.<locals>.<genexpr>m  r   r   )r   r   r   s   ``   r=   r*   r*   e  r   r@   prefix_tree	full_treec                 T   ^^ / mS[         S[        SS4UU4S jjn[        UU UTS9  T$ )a  Return a list of broadcasted leaves in ``prefix_tree`` to match the number of leaves in ``full_tree``.

If a ``prefix_tree`` is a prefix of a ``full_tree``, this means the ``full_tree`` can be
constructed by replacing the leaves of ``prefix_tree`` with appropriate **subtrees**.

This function returns a list of leaves with the same size as ``full_tree``. The leaves are
replicated from ``prefix_tree``. The number of replicas is determined by the corresponding
subtree in ``full_tree``.

>>> broadcast_prefix(1, [1, 2, 3])
[1, 1, 1]
>>> broadcast_prefix([1, 2, 3], [1, 2, 3])
[1, 2, 3]
>>> broadcast_prefix([1, 2, 3], [1, 2, 3, 4])
Traceback (most recent call last):
    ...
ValueError: list arity mismatch; expected: 3, got: 4; list: [1, 2, 3, 4].
>>> broadcast_prefix([1, 2, 3], [1, 2, (3, 4)])
[1, 2, 3, 3]
>>> broadcast_prefix([1, 2, 3], [1, 2, {"a": 3, "b": 4, "c": (None, 5)}])
[1, 2, 3, 3, 3, 3]

Args:
    prefix_tree (pytree): A pytree with the same structure as a prefix of ``full_tree``.
    full_tree (pytree): A pytree with the same structure as a suffix of ``prefix_tree``.
    is_leaf (callable, optional): An extra leaf predicate function that will be called at each
        flattening step. The function should have a single argument with signature
        ``is_leaf(node) -> bool``. If it returns :data:`True`, the whole subtree being treated
        as a leaf. Otherwise, the default pytree registry will be used to determine a node is a
        leaf or not. If the function is not specified, the default pytree registry will be used.

Returns:
    A list of leaves in ``prefix_tree`` broadcasted to match the number of leaves in ``full_tree``.
r   subtreer6   Nc                 X   > [        UTS9nTR                  U /UR                  -  5        g r   )r!   extend
num_leaves)r   r   subtreespecr]   results      r=   
add_leaves$broadcast_prefix.<locals>.add_leaves  s(    $Wg>qcK2223r@   rm   )r   r   r$   )r   r   r]   r   r   s     ` @r=   broadcast_prefixr   p  sG    N F4c 4F 4t 4 4 	 Mr@   c                     [        U5      (       d   e[        S/UR                  -  U5      n [        XUS9$ ! [         a     g f = f)Nr   rm   )r[   r   r   r   
ValueError)r\   re   r]   r   s       r=   _broadcast_to_and_flattenr     sR    
 #8,,,,sX%8%88(CIAA s   	8 
AAprotocolc                     [        U 5      (       d  [        S[        U 5       S35      e[        S/U R                  -  U 5      n[
        R                  " U5      n[
        R                  " X1S9$ )z&Serialize a treespec to a JSON string.z`treespec_dumps(treespec): Expected `treespec` to be instance of PyTreeSpec but got item of type rg   r   )r   )r[   rh   ri   r   r   rP   r!   r+   )re   r   
dummy_treeorig_treespecs       r=   r+   r+     sm    "8,,//3H~.>aA
 	

  h&9&9 98DJ!00<M''IIr@   
serializedc                     [         R                  " U 5      n[         R                  " S/UR                  -  U5      n[	        U5      nU$ )z*Deserialize a treespec from a JSON string.r   )rP   r,   r   r   r!   )r   r   r   re   s       r=   r,   r,     sH     "00<M--	
m&&&J j)HOr@   c                   "    \ rS rSrS\4S jrSrg)
_DummyLeafi  r6   c                     g)N*rw   )selfs    r=   __repr___DummyLeaf.__repr__  s    r@   rw   N)__name__
__module____qualname____firstlineno__strr   __static_attributes__rw   r@   r=   r   r     s    # r@   r   c                     [        [        U R                  5       Vs/ s H  n[        5       PM     snU 5      n[	        U5      $ s  snf r;   )r   ranger   r   repr)re   _r   s      r=   r-   r-     sB    $X%8%89:9!9:J 
 	;s   Ac                   &    \ rS rSrS\S\4S jrSrg)LeafSpecMetai  instancer6   c                 F    [        U5      =(       a    UR                  5       $ r;   )r[   r]   )r   r   s     r=   __instancecheck__LeafSpecMeta.__instancecheck__  s    &x0GX5E5E5GGr@   rw   N)r   r   r   r   objectr   r   r   rw   r@   r=   r   r     s    H& HT Hr@   r   c                       \ rS rSrSS jrSrg)r   i  c                 *    [         R                  " SS9$ )NT)r`   )rV   treespec_leaf)rH   s    r=   __new__LeafSpec.__new__  s    ##66r@   rw   N)r6   r   )r   r   r   r   r   r   rw   r@   r=   r   r     s    7r@   r   )	metaclassc                     [        S5      e)a  Flattens a pytree like :func:`tree_flatten`, but also returns each leaf's key path.

Args:
    tree: a pytree to flatten. If it contains a custom type, that type must be
        registered with an appropriate `tree_flatten_with_path_fn` when registered
        with :func:`register_pytree_node`.
    is_leaf: An extra leaf predicate function that will be called at each
        flattening step. The function should have a single argument with signature
        ``is_leaf(node) -> bool``. If it returns :data:`True`, the whole subtree being treated
        as a leaf. Otherwise, the default pytree registry will be used to determine a node is a
        leaf or not. If the function is not specified, the default pytree registry will be used.
Returns:
    A tuple where the first element is a list of (key path, leaf) pairs, and the
    second element is a :class:`TreeSpec` representing the structure of the flattened
    tree.
rL   rN   rb   s     r=   r   r     s    ( M
NNr@   c                     [        S5      e)a  Gets the leaves of a pytree like ``tree_leaves`` and returns each leaf's key path.

Args:
    tree: a pytree. If it contains a custom type, that type must be
        registered with an appropriate `tree_flatten_with_path_fn` when registered
        with :func:`register_pytree_node`.
    is_leaf: An extra leaf predicate function that will be called at each
        flattening step. The function should have a single argument with signature
        ``is_leaf(node) -> bool``. If it returns :data:`True`, the whole subtree being treated
        as a leaf. Otherwise, the default pytree registry will be used to determine a node is a
        leaf or not. If the function is not specified, the default pytree registry will be used.
Returns:
    A list of (key path, leaf) pairs.
rL   r   rb   s     r=   r    r      s    $ M
NNr@   c                    [        S5      e)af  Like :func:`tree_map`, but the provided callable takes an additional key path argument.

Args:
    func: A function that takes ``2 + len(rests)`` arguments, to be applied at the
        corresponding leaves of the pytrees. The first positional argument
        to ``func`` is the key path of the leaf in question. The second
        positional argument is the value of the leaf.
    tree: A pytree to be mapped over, with each leaf providing the first positional
        argument to function ``func``.
    rests: A tuple of pytrees, each of which has the same structure as
        ``tree`` or has ``tree`` as a prefix.
    is_leaf: An extra leaf predicate function that will be called at each
        flattening step. The function should have a single argument with signature
        ``is_leaf(node) -> bool``. If it returns :data:`True`, the whole subtree being treated
        as a leaf. Otherwise, the default pytree registry will be used to determine a node is a
        leaf or not. If the function is not specified, the default pytree registry will be used.

Returns
    A new pytree with the same structure as ``tree`` but with the value at each leaf given by
    ``func(keypath, x, *xs)`` where ``keypath`` is the key path at the
    corresponding leaf in ``tree``, ``x`` is the value at that leaf, and
    ``xs`` is the tuple of values at corresponding nodes in ``rests``.
rL   r   rp   s       r=   r#   r#     s    : M
NNr@   kpc                     [        S5      e)z9Given a key path, return a pretty-printed representation.rL   r   )r   s    r=   r   r   7      
M
NNr@   c                     [        S5      e)zAGiven an object and a key path, return the value at the key path.rL   r   )rX   r   s     r=   r   r   <  r   r@   rw   r;   )g__doc__rA   r   r   collections.abcr   typingr   r   r   r   r   r	   typing_extensionsr
   r   rV   torch._vendor.packaging.versionr   __version__ImportErrorr   r   torch.utils._pytreeutils_pytreerP   r   __all__dict_insertion_ordered__TORCH_DICT_SESSION	__enter__r1   r2   r3   r4   r   r   r   listr   r   OpTreeUnflattenFuncr   r   r   KeyPathFlattenWithKeysFuncrC   ri   r   r   FutureWarningrT   rO   r[   r   ra   r   r   r   r   r!   r"   r$   Type2Type3r   r   TypeAnyFn2Fn3FnFnAny	MapOnlyFnry   r%   r&   r'   r(   r)   r*   r   r   intr+   	lru_cacher,   r   r-   r   r   r   r    r#   r   r   _NODE_REGISTRY_LOCK_cxx_pytree_importedr8   r9   _cxx_pytree_pending_importsclearrw   r@   r=   <module>r     sM    
  $ D D 0  3 6!22
	Q 
  * + + 4 F 44TWM       CLCLCLCL 	xtCy''9!::;(3-169:#7?@ y/9:  /!2G!;< 
#
xtE(C-4H/I3/N)OOP  *=  +/9==A:><	c<<  <
 #3-< ""56< $$9:< ##67< 
<~ I +/9==A0	c00  0
 #3-0 ""560 $$9:0 
0
0p +/9==A
	c

  

 #3-
 ""56
 $$9:
 

0% %F8,< % 37#
#hx~./# 
#P 37(
(hx~./( 49h(V38C= 3H 3 3: 37 
 hx~./  c] J 37 
 hx~./  
#Y J 37 
 hx~./   N 37	.
38
.
. . hx~./	.
 .j 37	#
38
#
# # hx~./	#
 #L 	d1gtAwd1gtAwQ'(wDIuT#Y^4eooEFGDIuT#Y^445Gad}a aAg"#qc1f#aS(C5#:../	
 
DG 9R3Z3H  
 
E!Q$K yQ3Y7P  
 
E!Q'N )C1aQTDU:V  

 
G 9U3Z3H  
 
HcUD[$9 5QT:AV  
+ (C5$;*?!?@+uSz+\ 
 37 QV* 	
 hx~./7  
 
 37 aCi. 	
 hx~./ A;  
 
 37 aAsl
 	
 hx~./ Aq>  
 
 37 * 	
 hx~./"  
 
 37 * 	
 hx~./#SE4K0  
 37R *R 	R
 hx~./R (C5$;*?!?@R R 
 37 QV* 	
 hx~./7  
 
 37 aCi. 	
 hx~./ A;  
 
 37 aAsl
 	
 hx~./ Aq>  
 
 37 * 	
 hx~./"  
 
 37 * 	
 hx~./#SE4K0  
 37S *S 	S
 hx~./S (C5$;*?!?@S S 37%
C5$;
%
% hx~./% 
	% 37%
C5$;
%
% hx~./% 
	% 
 37 QW+ 	
 hx~./7 
 
 
 37 aDj/ 	
 hx~./A; 
 
 
 37 aAtm
 	
 hx~./Aq> 
 
 37K +K 	K
 hx~./KK 
K 
 37 QW+ 	
 hx~./7 
 
 
 37 aDj/ 	
 hx~./A; 
 
 
 37 aAtm
 	
 hx~./Aq> 
 
 37K +K 	K
 hx~./KK 
K 37333 hx~./3 
#Y	3B 37



 hx~./
 d3i	

JX 
J# 
J# 
J s x   
h 3 H4> H
7x< 7 37O
Ohx~./O 4gsl#$h./O2 37O
Ohx~./O 
%
O2 37	O
38
O
O O hx~./	O
 O@Ow O3 O
O O' Oc O
 &&)-M&rLD&%AAf%t6v6 B--335f '&&s    Aa00
a>