
    ^j|4                        U d Z ddlZddlmZ ddlmZ ddlmZmZm	Z	m
Z
 ddl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 ddlmZmZ d	d
lmZ d	dlmZ d	dlmZ d	dlm Z  ejB                  e"z  e#z  e$z  Z%de%fdZ&de%fdZ'e G d d             Z( G d d      Z)e G d d             Z*e*e+d   z  Z,e
e-d<    G d de      Z.y)aV  This file implements the IndexPropagation ops handler, which wraps an
underlying handler to add a limited form of constant propagation, as well as
propagation of sympy expressions downstream of ops.index_expr calls.

For example, say we have the IR:

    tmp0 = ops.index_expr(x, torch.int32)
    tmp1 = ops.constant(2, torch.int32)
    tmp2 = ops.mul(tmp0, tmp1)
    tmp3 = ops.indirect_indexing(tmp2, x_size)
    tmp4 = ops.load("buf0", tmp3)

The underlying handler would just see:

    ops.load("buf0", x * 2)

This is limited by the set of operators handled in the sympy expression
printers.

    N)Sequence)	dataclass)AnyLiteraloverload	TypeAlias)dtype_to_typeis_integer_dtype)FloorDivMaxMinModularIndexingWhere)bound_sympyValueRanges   )DefaultHandler)statically_known_true)generate_assert)Vvalc                     t        | t        j                        r| j                  S t        | t        t
        t        f      S N)
isinstancesympyBasic	is_numberintfloatboolr   s    l/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/torch/_inductor/index_propagation.py_is_constantr#   ,   s.    #u{{#}}cC-..    c                 d    t        | t        j                        rt        |       j                  S | S r   )r   r   Exprr   upperr!   s    r"   upper_boundr(   2   s%    %/UZZ%@;s!!IcIr$   c                   F    e Zd ZU dZeed<   ej                  ed<   d Zd Z	y)	TypedExprz'A SymPy expression with associated typeexprdtypec                 ,    t        | j                        S r   )r#   r+   selfs    r"   is_constantzTypedExpr.is_constant=   s    DII&&r$   c                    t        | j                        r| j                  }t        |t        j                        r|j                  d      } t        | j                        |      }t        | j                        rst        j                  | j                        j                  }| j                  j                  r|d|dz
  z  z   }|d|z  z  }| j                  j                  r|d|dz
  z  z
  }|| _        y y )NT)identity   r   )r#   r+   r   r   r&   expandr	   r,   r
   torchiinfobits	is_signed)r/   r+   r7   s      r"   __post_init__zTypedExpr.__post_init__@   s    		"99D$

+{{D{1,=,T2D

+{{4::.33::''!q/1Dag~::''!q/1DDI #r$   N)
__name__
__module____qualname____doc__	_ExprType__annotations__r5   r,   r0   r9    r$   r"   r*   r*   6   s    1
O;;'r$   r*   c                      e Zd ZdZededefd       Zedeez  e	z  de
j                  defd       Zedej                  ez  de
j                  defd       Zedej                  ez  de
j                  defd       Ze	 	 ddede
j                  d
e
j                  d	z  de	def
d       Zededefd       Zededefd       Zedededefd       Zedededefd       Zedededefd       Zededefd       Zedededefd       Zedededed	z  fd       Zedededed	z  fd       Zedededefd       Zedededefd       Zy	)SymPyOpszAn ops handler where all IR values are SymPy expressions

    When a value cannot be represented as a SymPy expression, the method is
    either not defined, or returns NotImplemented

    valuereturnc                     | S r   r@   )rC   s    r"   r2   zSymPyOps.identityX   s    r$   r,   c                     t        | |      S r   r*   rC   r,   s     r"   constantzSymPyOps.constant\       &&r$   c                     t        | |      S r   rG   rH   s     r"   
index_exprzSymPyOps.index_expr`   rJ   r$   c                     t        | |      S r   rG   rH   s     r"   
value_exprzSymPyOps.value_exprd   rJ   r$   N	src_dtypeuse_compute_typesc                 .    t        | j                  |      S r   )r*   r+   )rC   r,   rO   rP   s       r"   to_dtypezSymPyOps.to_dtypeh   s     U++r$   xc                 T    t        t        | j                        | j                        S r   )r*   absr+   r,   rS   s    r"   rU   zSymPyOps.absq   s    QVVagg..r$   c                 \    t        | j                  | j                  z  | j                        S r   r*   r+   r,   rV   s    r"   squarezSymPyOps.squareu   s    !&&!''22r$   yc                     t        j                  | j                  |j                        }t        | j                  |j                  z   |      S r   r5   promote_typesr,   r*   r+   rS   rZ   result_types      r"   addzSymPyOps.addy   5    ))!''177;!&&+66r$   c                     t        j                  | j                  |j                        }t        | j                  |j                  z
  |      S r   r\   r^   s      r"   subzSymPyOps.sub~   ra   r$   c                     t        j                  | j                  |j                        }t        | j                  |j                  z  |      S r   r\   r^   s      r"   mulzSymPyOps.mul   ra   r$   c                 D    t        | j                   | j                        S r   rX   rV   s    r"   negzSymPyOps.neg   s    !&&!''**r$   c                     t        j                  | j                  |j                        }t        |      st        S t        t        | j                  |j                        |      S r   )r5   r]   r,   r
   NotImplementedr*   r   r+   r^   s      r"   floordivzSymPyOps.floordiv   sF    ))!''177;,!!!&&!&&1;??r$   c                    t        j                  | j                  |j                        }t        |      st        S t        | j                  t        j                  j                  |j                        }t        ||      S r   )r5   r]   r,   r
   ri   r   r+   r   SOner*   )rS   rZ   r_   result_exprs       r"   modzSymPyOps.mod   sU    ))!''177;,!!%affeggkk166Bk22r$   c                    t        j                  | j                  |j                        }t        |      st        S t        j                  | j                        }t        j                  |j                        }|j                  ^|j                  |j                  k(  rEt        | j                  t
        j                  j                  |j                        }t        ||      S t        S r   )r5   r]   r,   r
   ri   r   sympifyr+   is_nonnegativeis_positiver   rl   rm   r*   )rS   rZ   r_   x_expry_exprrn   s         r"   	remainderzSymPyOps.remainder   s    ))!''177;,!!qvv&qvv& !!-%%););;)!&&%''++qvvFK[+66r$   c                     t        j                  | j                  |j                        }t        t	        | j
                  |j
                        |      S r   )r5   r]   r,   r*   r   r+   r^   s      r"   minimumzSymPyOps.minimum   8    ))!''177;QVVQVV,k::r$   c                     t        j                  | j                  |j                        }t        t	        | j
                  |j
                        |      S r   )r5   r]   r,   r*   r   r+   r^   s      r"   maximumzSymPyOps.maximum   ry   r$   )NF)r:   r;   r<   r=   staticmethodr   r2   r   r   r    r5   r,   r*   rI   r   r&   rL   rN   rR   rU   rY   r`   rc   re   rg   rj   ro   rv   rx   r{   r@   r$   r"   rB   rB   P   s        'ed* '5;; '9 ' ' '%**s* '5;; '9 ' ' '%**s* '5;; '9 ' '  )-"'	,,{{, ;;%,  	,
 
, , /y /Y / / 3) 3	 3 3 7y 7Y 79 7 7 7y 7Y 79 7 7 7y 7Y 79 7 7 +y +Y + + @I @) @	 @ @ 3y 3Y 39t+; 3 3 Y 9 T1A  " ;9 ; ;y ; ; ;9 ; ;y ; ;r$   rB   c                   F    e Zd ZU eed<   dZeed<   ededd fd       Z	d Z
y)	IndexPropVarrC   Fis_symbolicr+   rD   c                     t        | d      S )NTr   )r~   )r+   s    r"   new_symboliczIndexPropVar.new_symbolic   s    Dd33r$   c                 `    | j                   r"t        | j                  t              sJ d       y y )Nz.Symbolic IndexPropVar must contain a TypedExpr)r   r   rC   r*   r.   s    r"   r9   zIndexPropVar.__post_init__   s.    ##z$**i'H 	
<	
H'H#r$   N)r:   r;   r<   r   r?   r   r    r|   r*   r   r9   r@   r$   r"   r~   r~      s6    JK49 4 4 4
r$   r~   )IndexPropResult.r   c            	       6   e Zd ZdZdedeej                  ej                  f   deej                  ej                  f   ddfdZ	dej                  d	e
j                  defd
Zdej                  d	e
j                  defdZdeez  defdZdefdZeded   dee   deeef   defd       Zededee   deeef   defd       Zdedee   deeef   defdZdedee   deeef   defdZdedeedf   deeef   defdZd Z	 	 ddeez  dededefdZy)IndexPropagationzOps wrapper that tries to propagate constant and index_expr values through the computation.

    This aims to maximize the compile time simplification possible, and convert
    indirect indexing from arange into normal static indexing.

    inneriter_rangesindirect_var_rangesrD   Nc                 r   || _         t        j                  j                  j                  | _        |j                         D ci c]  \  }}|t        dt        |      dz
          }}}t        t        j                  | j                  j                  j                         |j                                     | _        || _        g }|j                         D ]-  \  }}	|j                  d|k         |j                  ||	k         / t        |      | j                  j                         z   | _        y c c}}w )Nr   r   )_innerr   graphsizevars	shape_envitemsr   r(   tuple	itertoolschainvar_to_ranger   append
get_axiomsaxioms)
r/   r   r   r   kvr   r   rS   ss
             r"   __init__zIndexPropagation.__init__   s	    ))33 ?J>O>O>Q
6:aA{1k!nq011
 
 "OODNN77==?ASASAUV

 $7 %%' 	!DAqMM!q&!MM!a% 	! Fmdnn&?&?&AA
s   #D3r+   r,   c                     t        |      r- t        |      |      }| j                  j                  ||      S | j                  j	                  ||      S r   )r#   r	   r   rI   rL   )r/   r+   r,   r   s       r"   materialize_exprz!IndexPropagation.materialize_expr   sI    &-&t,C;;''U33{{%%dE22r$   c                 X    | j                  | j                  j                  ||            S r   )wrapr   rN   )r/   r+   r,   s      r"   rN   zIndexPropagation.value_expr   s"    yy//e<==r$   ac                      t        |t        t        f      rt         fd|D              S t        |t              s|S |j                  r: j                  |j                  j                  |j                  j                        S |j                  S )Nc              3   @   K   | ]  }j                  |        y wr   )unwrap.0r   r/   s     r"   	<genexpr>z*IndexPropagation.unwrap.<locals>.<genexpr>   s     3AQ3   )	r   listr   r~   r   r   rC   r+   r,   r/   r   s   ` r"   r   zIndexPropagation.unwrap   sf    a$'3333!\*H ==((qww}}EEwwr$   c                 n     t        |t        t        f      rt         fd|D              S t        |      S )Nc              3   @   K   | ]  }j                  |        y wr   )r   r   s     r"   r   z(IndexPropagation.wrap.<locals>.<genexpr>  s     1!11r   )r   r   r   r~   r   s   ` r"   r   zIndexPropagation.wrap  s,    a$'1q111Ar$   nameindirect_indexingargskwargsc                      y r   r@   r/   r   r   r   s       r"   fallbackzIndexPropagation.fallback	  s     r$   c                      y r   r@   r   s       r"   r   zIndexPropagation.fallback  s     r$   c                    |D cg c]  }| j                  |       }}|j                         D ci c]  \  }}|| j                  |       }}}| j                   t        | j                  |      |i |      S c c}w c c}}w r   )r   r   r   getattrr   )	r/   r   r   r   r   new_argsr   r   
new_kwargss	            r"   r   zIndexPropagation.fallback  sx     -11qDKKN114:LLNCDAqaQ'C
Cyy3d3XLLMM 2Cs
   A<Bc                    dt         t        z  dt         fd}|D cg c]
  } ||       }}|j                         D ci c]  \  }}| ||       }	}} t        t        |      |i |	}
|
t
        uxr( |
j                         xs |
j                  j                  }|s| j                  |||      S t        j                  |
      S c c}w c c}}w )Nr   rD   c                 >    t        | t              s| S | j                  S r   )r   r~   rC   )r   s    r"   r   z0IndexPropagation.propagate_sympy.<locals>.unwrap"  s    a.77Nr$   )r   r~   r   r   rB   ri   r0   r+   
is_integerr   r   )r/   r   r   r   r   r   r   r   r   r   new_expris_valid_exprs               r"   propagate_sympyz IndexPropagation.propagate_sympy  s    	cL( 	S 	
 (,,!F1I,,/5||~>tq!al>
>*78T*HC
C 6 
   ">hmm&>&> 	
 ==tV44((22 ->s   C C.c                 D   t        t        |      s| j                  |||      S t        j                  ||j                               D cg c]  }t        |t              r| }}t        d |D              s| j                  |||      S | j                  |||      S c c}w )Nc              3   4   K   | ]  }|j                     y wr   r   )r   r   s     r"   r   z,IndexPropagation._default.<locals>.<genexpr><  s     8Q1==8s   )
hasattrrB   r   r   r   valuesr   r~   allr   )r/   r   r   r   r   var_argumentss         r"   _defaultzIndexPropagation._default3  s    x&==tV44 __T6==?;
!\* 
 

 8-88==tV44##D$77
s   
Bc                     g | j                   d | j                  j                         D        }t        | j                  || j
                  |      S )a  
        Given some iter_ranges, return a function that given an expression, returns whether
        it is true or false using value ranges, guard knowledge and runtime_asserts.

        FIXME I think this may not be entirely right, as we may not be able to use all runtime_asserts
              If this is an issue, just use guards in `self.axioms`.

              The proper way of handling this would be to have a global shape_env that adds
              runtime_asserts as they happen in the code. Then, it should be used in SimplifyIndexing
              to perform wrap_expr and in CSEProxy.check_bounds to elide upper / lower bounds also
              for indirect_indexing
        c              3   V   K   | ]!  \  }}|t        d t        |      dz
        f # yw)r   r   N)r   r(   )r   r   r   s      r"   r   z3IndexPropagation.statically_true.<locals>.<genexpr>P  s1      Aq K;q>A#567s   '))r   r   r   r   r   r   )r/   er   s      r"   statically_truez IndexPropagation.statically_trueA  sS    

 44::<
 %T^^Q\RRr$   indexsizecheckc                     t        |t              r|j                  rt        j                  |j
                  j                        } fd} j                  d|k        xs |xr  j                   |k        } j                  |k        }|r ||      }t        |      r" j                  d|ft        | |              |S  j                  d|||fi       j
                  }	|	S )Nc                     j                  d| k        r| S j                  | dk        r| z   S t        | dk  | z   |       S )Nr   )r   r   )r+   r/   r   s    r"   	wrap_exprz5IndexPropagation.indirect_indexing.<locals>.wrap_exprh  sM    ''T	2K))$(3$;& 4$;==r$   r   check_bounds)lowerr'   r   )r   r~   r   r   rq   rC   r+   r   r   r   dict)
r/   r   r   r   wrap_negr+   r   can_prove_lowercan_prove_upperindirect_vars
   ` `       r"   r   z"IndexPropagation.indirect_indexingX  s     e\*u/@/@ ==!1!12D> #2219= @T114%4-@  #224$;?O u%"4L?2o:MN
 K}}%uh!?

% 	 r$   )TT)r:   r;   r<   r=   r   r   r   Symbolr&   r   r5   r,   r   r   rN   r~   r   r   r   r   r   strr   r   r   r   r   r    r   r@   r$   r"   r   r      s&   BB %,,

23B "%,,

":;	B
 
B23UZZ 3 3 3>uzz >%++ >/ >l* s  
 )* sm S#X	
 
  '}6:38n	 NN'}N6:38nN	N33'}36:38n3	3*8S 8c3h 8c3h 8TW 8S6 +\!+ + 	+ 
+r$   r   )/r=   r   collections.abcr   dataclassesr   typingr   r   r   r   r   r5   torch._prims_commonr	   r
   torch.utils._sympy.functionsr   r   r   r   r   torch.utils._sympy.value_rangesr   r   ops_handlerr   r   r   utilsr   virtualizedr   r&   r   r   r    r>   r#   r(   r*   rB   r~   r   r   r?   r   r@   r$   r"   <module>r      s   *  $ ! 4 4   ? S S D ' + "  JJ$t+	/i /JY J   2g; g;T 
 
 
 *E2H,II Iw~ wr$   