+
    &j                         R t ^ RIHt ^ RIt^ RIHtHt ^ RIt^ RIH	t	 ^RI
Ht ^RIHt ^RIHt ]'       d   ^ RIHt R	 R
 lt ! R R4      tRR R lltRR R lltR R ltRR R lltRR R lltR# )zM
Python implementation of function wrapping functionality for functorch.dim.
)annotationsN)AnyTYPE_CHECKING)tree_map)DimEntry)EnableAllLayers)
TensorInfo)Callablec                    V ^8  d   QhRRRR/# )   tensorztorch.Tensorreturn )formats   "k/Users/jameslopez/projects/CWCArchive/cwc-podcast/.venv/lib/python3.14/site-packages/functorch/dim/_wrap.py__annotate__r      s      |      c                    V # )z8Handle tensor conversion for torch function integration.r   )r   s   &r   handle_from_tensorr      s    Mr   c                  6    ] tR t^tRtRR R lltR R ltRtR# )	WrappedOperatorzH
This class wraps PyTorch operations to support first-class dimensions.
c               $    V ^8  d   QhRRRRRR/# )r   origr	   wrapper_implementationdim_namestrr   )r   s   "r   r   WrappedOperator.__annotate__    s)     v vv6>vJMvr   c                	J   Wn         W n        \        VR R4      V n        \        VRR4      V n        W0n        RV n        ^ V n        ^V n        RV n	        RV n
        V P                  '       d8   V P
                  '       d$   V P                   RV P
                   R2V n        R# R# R# )__name__ __doc__NFTz
Argument 'z5' can be either an integer or a torchdim.Dim object.
)r   r   getattrnamedocr   is_pointwise
dim_offsetkeepdim_offset
single_dimreduce)selfr   r   r   s   &&&&r   __init__WrappedOperator.__init__    s     	&<#D*b1	4D1 ! 888((<>tuDH &8r   c                   V ^8  d   QhRR/# )r   r   r	   r   )r   s   "r   r   r   3   s      ( r   c                   a  R V 3R llp\         P                  ! VS P                  RRR7       S P                  Vn        V# )z@Create a wrapped function that calls our wrapper implementation.c               $    V ^8  d   QhRRRRRR/# )r   argsr   kwargsr   r   )r   s   "r   r   .WrappedOperator.function.<locals>.__annotate__6   s&     	F 	F 	Fs 	Fs 	Fr   c                 0   < SP                   ! S.V O5/ VB # )N)r   )r/   r0   r)   s   *,r   wrapped_func.WrappedOperator.function.<locals>.wrapped_func6   s    ..tEdEfEEr   )assignedupdated)r   r   )	functoolsupdate_wrapperr   r#   r    )r)   r3   s   f r   functionWrappedOperator.function3   s@    	F 	F 	  $))mR	
  $xxr   )
r   r%   r#   r$   r&   r"   r   r(   r'   r   N)dim)r   
__module____qualname____firstlineno__r    r*   r9   __static_attributes__r   r   r   r   r      s    v& r   r   c               (    V ^8  d   QhRRRRRRRR/# )	r   r;   r   ndimintkeepdimboolr   r   r   )r   s   "r   r   r   B   s(      3 c D X r   c                    ^RI Hp \        W4      '       d   V'       d   \        R4      h\	        V 4      # \        V \
        4      '       d   T pV^ 8  d   WA,          pK  \	        V4      # \	        4       # )z:Convert single dimension specification to DimEntry object.)Dimz8cannot preserve first-class dimensions with keepdim=True)r   rF   
isinstance
ValueErrorr   rB   )r;   rA   rC   rF   is   &&&  r   	_wrap_dimrJ   B   s[    #WXX}	C		1fIA{zr   c               (    V ^8  d   QhRRRRRRRR/# )	r   r;   r   rA   rB   rC   rD   r   zlist[DimEntry]r   )r   s   "r   r   r   S   s(     	 	C 	s 	T 	n 	r   c                    \        WV4      p. pVP                  4       '       g   VP                  V4       V# V  F  pVP                  \        WQV4      4       K   	  V# )z<Convert dimension specification to list of DimEntry objects.)rJ   is_noneappend)r;   rA   rC   deresultds   &&&   r   
_wrap_dimsrR   S   sU    	3g	&BF::<<b M AMM)AW56 Mr   c               (    V ^8  d   QhRRRRRRRR/# )r   wrapperr   r/   r   r0   r   r   )r   s   "r   r   r   _   s.     w) w) w) w)s w)s w)r   c                V	  aa V'       g   \        R4      hVP                  V P                  4      pVfF   V P                  \	        V4      8  d,   V P                  ^,           pV\	        V4      8  d	   W,          pVf   \
        P                  ! V^ ,          RRR7      oS'       g   V P                  ! V/ VB # \        SP                  4      ;_uu_ 4       pSP                  f   \        R4      hVP                  SP                  SP                  4       \        V4      p\        SP                  4      V^ &   V P                  ! V/ VB pVP                  VSP                   4      uuRRR4       # \
        P                  ! V^ ,          4      oS'       g   V P                  ! V/ VB # RpV P"                  '       dj   VP                  R4      p	V	fF   V P$                  \	        V4      8  d,   V P$                  ^,           p
V
\	        V4      8  d	   W,          p	V	e   \'        V	4      pSP)                  4       p\+        W;V4      p. pR.\	        SP                  4      ,          pV F  pRp\-        SP                  4       F  w  ppVV8X  g   K  Tp M	  Vf   \-        SP                  4       F5  w  pp\/        VR4      '       g   K  VP1                  V4      '       g   K3  Tp M	  Vf6   SP                   Uu. uF  p\3        V4      NK  	  pp\        R	V R
V 24      hRVV&   VP5                  V4       K  	  . oV P"                  '       dK   V'       gC   \-        SP                  4       F(  w  ppVV,          '       d   K  SP5                  V4       K*  	  MSP                  R,          o\	        V4      ^8X  d   V^ ,          pM\7        V4      p\        V4      pVP9                  4       pSP:                  f   \        R4      h\        SP:                  4      V^ &   V P                  V9   d   VVV P                  &   M2V P                  ^,           pV\	        V4      8  d   \        V4      pVWd&   V P                  ! V/ VB pR VV3R llp\=        VV4      #   + '       g   i     EL3; iu upi )zB
This is the core method that handles dimension-aware operations.
z%Expected at least one argument (self)NTF)ensure_batchedensure_presentz%Expected batchedtensor to be non-NonerC   matcheszTensor with dimensions z does not contain :NNNzExpected tensor to be non-Nonec                    V ^8  d   QhRRRR/# )r   objr   r   r   )r   s   "r   r   (patched_dim_method.<locals>.__annotate__   s        r   c                   < \        V \        P                  4      '       d$   ^RIHp VP	                  V SSP
                  4      # V # )   )Tensor)rG   torchr^   r   from_positional
has_device)rZ   r^   info
new_levelss   & r   wrap_result'patched_dim_method.<locals>.wrap_result   s5    c5<<(( ))#z4??KK
r   )rH   getr   r%   lenr   creater   r   levelsbatchedtensorAssertionErrorinplace_update_layerslistr   from_batchedra   r(   r&   rD   rA   rR   	enumeratehasattrrX   r   rN   tuplecopyr   r   )rT   r/   r0   dim_argdim_idxguardnew_argsrP   rC   keepdim_argkeepdim_idxrA   dimsdim_indicesseenrQ   midxrI   level
level_strs
py_indices
new_kwargsrd   rb   rc   s   &*,                    @@r   patched_dim_methodr   _   s     @AA jj))*G7--D	9$$q(SYmG   aeT<<000T[[))U!!)$%LMM''(:(:DKKHDzH,T-?-?@HQK\\86v6F%%fdoo> *) T!W%D||T,V,, G~~~jj+7#9#9CI#E!0014KSY&"/";'G 99;DgW-D  K7S%%D!$++.HAuz /
 <%dkk255),,q1A1AD 3
 |6:kkBkUc%jk
B -j\9KA3O  T
4 + 0 J~~~g!$++.HAu77!!%( / [[^
 ;1%a.
;'
 DzHJ{{=>>$T[[1HQK :%'1
7##$$$q(S]"H~H *H \\82z2F  K((E *))` Cs   BR4R&R#	c               4    V ^8  d   QhRRRRRRRRRR	R
R	RR/# )r   r   r	   r%   z
int | Noner&   r   z
str | Noner'   zbool | Noner(   r   r   )r   s   "r   r   r      sN        
     	 
      r   c                    T;'       g    Rp\        V \        V4      pVe   Wn        Ve   W&n        Ve   WFn        Ve   WVn        VP                  4       # )a  
Wrap a PyTorch function to support first-class dimensions.

Args:
    orig: Original function to wrap
    dim_offset: Offset for dimension argument (default: 0)
    keepdim_offset: Offset for keepdim argument (default: 1)
    dim_name: Name of dimension parameter (default: "dim")
    single_dim: Whether function takes single dimension (default: False)
    reduce: Whether function reduces dimensions (default: True)
r;   )r   r   r%   r&   r'   r(   r9   )r   r%   r&   r   r'   r(   rT   s   &&&&&& r   _wrapr      s^    &   5Hd$6AG'!!/'r   c               0    V ^8  d   QhRRRRRRRRRR	R
R/# )r   rT   r   funcr	   typesrq   r/   r0   zdict | Noner   r   r   )r   s   "r   r   r      sL     A AA
A A 	A
 A 	Ar   c                >    Vf   / p^RI Hp VP                  WW44      # )z8
Handle __torch_function__ calls for wrapped operators.
)_Tensor)r   r   __torch_function__)rT   r   r   r/   r0   r   s   &&&&& r   call_torch_functionr      s(     ~  %%d4@@r   )F)NNNNN)r   N)r    
__future__r   r7   typingr   r   r_   torch.utils._pytreer   
_dim_entryr   _enable_all_layersr   _tensor_infor   collections.abcr	   r   r   rJ   rR   r   r   r   r   r   r   <module>r      s_    #  %  (   / $ (
$ $N"	w)t FA Ar   