
    ^j                         U d dl mZ d dlmZ d dlmZ ddlmZ ddlm	Z	 ddl
mZ g dZ G d	 d
      Z e       Zeed<    G d d      Z G d d      Z G d d      Zdedededef   dz  fdZy)    )Callable)Any)
OrderedSet   )EffectHolder)FakeImplHolder)RegistrationHandle)SimpleLibraryRegistrySimpleOperatorEntry	singletonSymmMemArgsHolderc                   8    e Zd ZdZd
dZdeddfdZdeddfd	Zy)r
   aJ  Registry for the "simple" torch.library APIs

    The "simple" torch.library APIs are a higher-level API on top of the
    raw PyTorch DispatchKey registration APIs that includes:
    - fake impl

    Registrations for these APIs do not go into the PyTorch dispatcher's
    table because they may not directly involve a DispatchKey. For example,
    the fake impl is a Python function that gets invoked by FakeTensor.
    Instead, we manage them here.

    SimpleLibraryRegistry is a mapping from a fully qualified operator name
    (including the overload) to SimpleOperatorEntry.
    returnNc                     i | _         y N_dataselfs    i/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/torch/_library/simple_registry.py__init__zSimpleLibraryRegistry.__init__#   s	    57
    qualnamer   c                 v    | j                   j                  |d       }|t        |      x| j                   |<   }|S r   )r   getr   )r   r   ress      r   findzSimpleLibraryRegistry.find&   s8    jjnnXt,;)<X)FFDJJx 3
r   zSimpleOperatorEntry | Nonec                 :    | j                   j                  |d       S r   r   r   r   r   s     r   r   zSimpleLibraryRegistry.get,   s    zz~~h--r   r   N)__name__
__module____qualname____doc__r   strr   r    r   r   r
   r
      s4    8S %: .C .$@ .r   r
   r   c                   6    e Zd ZdZdeddfdZedefd       Zy)r   zThis is 1:1 to an operator overload.

    The fields of SimpleOperatorEntry are Holders where kernels can be
    registered to.
    r   r   Nc                     || _         t        |      | _        t        |      | _        t        |      | _        t        |      | _        y r   )	r   r   	fake_implGenericTorchDispatchRuleHoldertorch_dispatch_rulesr   effectr   symm_mem_argsr    s     r   r   zSimpleOperatorEntry.__init__:   s@    %)7)A*84 	! %1$:0A(0Kr   c                     | j                   S r   )r*   r   s    r   abstract_implz!SimpleOperatorEntry.abstract_implE   s    ~~r   )	r"   r#   r$   r%   r&   r   propertyr   r0   r'   r   r   r   r   3   s8    L L L ~  r   r   c                   Z    e Zd ZdeddfdZdededef   defdZ	dededef   dz  fd	Z
y)
r+   r   r   Nc                      i | _         || _        y r   )r   r   r    s     r   r   z'GenericTorchDispatchRuleHolder.__init__K   s    57
%r   torch_dispatch_classfunc.c                       j                        rt         d j                         | j                  <   d fd}t	        |      S )Nz8 already has a `__torch_dispatch__` rule registered for c                        j                   = y r   r   r   r4   s   r   
deregisterz;GenericTorchDispatchRuleHolder.register.<locals>.deregisterX   s    

/0r   r!   )r   RuntimeErrorr   r   r	   )r   r4   r5   r9   s   ``  r   registerz'GenericTorchDispatchRuleHolder.registerO   sZ     99)*'((`aeanan`op  ,0

'(	1 "*--r   c                 :    | j                   j                  |d       S r   r   r8   s     r   r   z#GenericTorchDispatchRuleHolder.find]   s    zz~~2D99r   )r"   r#   r$   r&   r   typer   r   r	   r;   r   r'   r   r   r+   r+   J   s\    & & &.$(.08c0B.	.: :(382Dt2K :r   r+   c                   j    e Zd ZdZdeddfdZdee   defdZde	e   dz  fdZ
defd	Zd
edefdZy)r   zTracks which arguments of an operator require symmetric memory allocation.

    Used by Inductor during lowering to automatically realize tensors as
    symmetric memory buffers.
    r   r   Nc                      d | _         || _        y r   )_symm_mem_argsr   r    s     r   r   zSymmMemArgsHolder.__init__h   s    6:%r   	arg_namesc                     |st        d j                          j                  Add l}|j	                  t
              }|j                  d j                   j                  |       t        |       _        d fd}t        |      S )Nz)Cannot register empty arg_names list for r   z>Overwriting symm_mem arg registration for %s. Old: %s, New: %sc                      d  _         y r   r@   r   s   r   r9   z.SymmMemArgsHolder.register.<locals>.deregister   s    "&Dr   r!   )	
ValueErrorr   r@   logging	getLoggerr"   warningr   r	   )r   rA   rF   logr9   s   `    r   r;   zSymmMemArgsHolder.registerl   s    ;DMM?K  *##H-CKKP##	 )3	' "*--r   c                     | j                   S r   rD   r   s    r   r   zSymmMemArgsHolder.get   s    """r   c                     | j                   d uS r   rD   r   s    r   is_registeredzSymmMemArgsHolder.is_registered   s    ""$..r   arg_namec                 >    | j                   d uxr || j                   v S r   rD   )r   rM   s     r   is_symm_mem_argz!SymmMemArgsHolder.is_symm_mem_arg   s#    ""$.R8t?R?R3RRr   )r"   r#   r$   r%   r&   r   listr	   r;   r   r   boolrL   rO   r'   r   r   r   r   a   si    & & &.$s) .0B .0#Z_t+ #/t /S S Sr   r   opr4   r   .Nc                 r    t         j                  | j                        j                  j                  |      S r   )r   r   r$   r,   )rR   r4   s     r   find_torch_dispatch_rulerT      s-     >>"//*??DD r   )collections.abcr   typingr   torch.utils._ordered_setr   effectsr   r*   r   utilsr	   __all__r
   r   __annotations__r   r+   r   r=   rT   r'   r   r   <module>r\      s    $  / ! % %. .: $9#:	  : .: :.*S *SZ#'c3h$r   