+
    &j              	         a  RP ta0 t ^ RIHt ^ RIt^ RIt^ RIHtHtHt ^RI	H
t
 ^ RIHtHt ^ RIHt ^ RIHt ^ RIHt ^ R	IHt ^ R
IHt ^ RIHt ^ RIHt ^ RItRR.t]! R4      t]! R4      t]P>                  ! ] 4      t! ^ RI"H#t$ ]PP                  PR                  t)R t*/ t+] ^ k R t,RRR R llt-]-! ])P\                  4      RR/R R ll4       t/]-! ])P`                  4      RSR R ll4       t1]-! ])Pd                  4      RSR R  ll4       t3]-! ])Ph                  4      RSR! R" ll4       t5]-! ])Pl                  4      RTR# R$ ll4       t7RRR% R& llt8]-! ])Pr                  ])Pt                  ])Pv                  ])Px                  ])Pz                  .4      RR/R' R( ll4       t>]-! ])P~                  4      R) R* l4       t@R+ tA]-! ])P                  ])P                  ])P                  .4      RR/R, R- ll4       tER. tFR/R/R0 R1 lltGR/R/R2 R3 lltH]-! ])P                  RR47      RR/R5 R6 ll4       tJ]-! ])P                  RR47      R7 R8 l4       tLR9 tM]-! ])P                  ])P                  ])P                  .4      RR/R: R; ll4       tQ]-! ])P                  RR47      R< R= l4       tS]-! ])P                  RR47      R> R? l4       tUR@R/RA RB lltVR@R/RC RD lltWR@R/RE RF lltX/ ])P\                  ]/b])P`                  ]1b])Pd                  ]3b])Ph                  ]5b])Pl                  ]7b])Pr                  ]>b])Pt                  ]>b])Pv                  ]>b])Pz                  ]>b])Px                  ]>b])P~                  ]@b])P                  ]Eb])P                  ]Eb])P                  ]Eb])P                  ]Qb])P                  ]Qb])P                  ]Qb])P                  ]J])P                  ]L])P                  ]S])P                  ]U/Ct+RG tY. RUOtZRH t[RI t\RJ RK lt]RL t^ ! RM R4      t_ ! RN RO]4      t`R#   ]% dN    ]&;QJ d    R RQ 4       F  '       g   K   RM	  RM! R RQ 4       4      '       d   ]!PO                  R4       ]t$ ELi ; i)V    )NoneTypeN)tree_maptree_flattentree_unflatten)ModuleTracker)AnyTypeVar)Callable)Iterator)	ParamSpec)defaultdict)TorchDispatchModeprodwrapsFlopCounterModeregister_flop_formula_T_PJITFunctionc              #   \   "   T F"  p\        \        P                  VR 4      R Jx  K$  	  R # 5iN)getattrtorchversion).0attrs   & p/Users/jameslopez/projects/CWCArchive/cwc-podcast/.venv/lib/python3.14/site-packages/torch/utils/flop_counter.py	<genexpr>r!      s$     
]F\d75==$-T9F\s   *,TFz@triton not found; flop counting will not work for triton kernelsc                 ^    \        V \        P                  4      '       d   V P                  # V # r   )
isinstancer   Tensorshape)is   &r    	get_shaper'   #   s!    !U\\""wwH    c                 8   a  \        S 4      R R/V 3R ll4       pV# )out_valNc                 B   < \        \        WV 34      w  rpS! VR V/VB # )	out_shape)r   r'   )r*   argskwargsr,   fs   $*, r    nfshape_wrapper.<locals>.nf+   s.    "*9tW6M"Ni$6)6v66r(   r   r/   r0   s   f r    shape_wrapperr3   *   s'    
1X7$ 7 7 Ir(   c                    V ^8  d   QhR\         \         \        \        3,          .\         \        \        3,          3,          /#    returnr
   r   r   )formats   "r    __annotate__r:   1   s6      XxB?O>PRZ[]_a[aRb>b5c r(   c                    a a R  VV 3R llpV# )c                t    V ^8  d   QhR\         \        \        3,          R\         \        \        3,          /# )r6   flop_formular7   r8   )r9   s   "r    r:   +register_flop_formula.<locals>.__annotate__3   s,      8BF#3 R8H r(   c                    <a  S'       g   \        S 4      o R  V 3R llp\        P                  P                  P	                  VS4       S # )c                    V ^8  d   QhRR/# )r6   r7   N )r9   s   "r    r:   Aregister_flop_formula.<locals>.register_fun.<locals>.__annotate__7   s     	1 	1 	1r(   c                    < \        V \        P                  P                  \        34      '       g   \        R V  R\        V 4       24      hV \        9   d   \        RV  24      hS\        V &   R# )z|register_flop_formula(targets): expected each target to be OpOverloadPacket (i.e. torch.ops.mylib.foo), or JitFunction, got z which is of type zduplicate registrations for N)	r#   r   _opsOpOverloadPacket_JITFunction
ValueErrortypeflop_registryRuntimeError)targetr=   s   &r    register=register_flop_formula.<locals>.register_fun.<locals>.register7   sp    v

(C(C\'RSS #H$6tF|nFG G &"%A&#JKK$0M&!r(   )r3   r   utils_pytree	tree_map_)r=   rL   get_rawtargetss   f r    register_fun+register_flop_formula.<locals>.register_fun3   s<    (6L	1 	1 	%%h8r(   rA   )rR   rQ   rS   s   ff r    r   r   1   s     & r(   r,   c                $    V ^8  d   QhR\         /# r5   int)r9   s   "r    r:   r:   I   s     	 	# 	r(   c               l    V w  rVVw  rxWg8w  d   \        RV RV 24      hWX,          ^,          V,          # )zCount flops for matmul.z3matmul: inner dimensions must match (k == k2), got  and AssertionError)	a_shapeb_shaper,   r-   r.   mkk2ns	   &&$*,    r    mm_floprb   H   sE    
 DAEBwRSTRUUZ[]Z^_``519q=r(   c                $    V ^8  d   QhR\         /# r5   rV   )r9   s   "r    r:   r:   U   s     % %# %r(   c                    \        W4      # )zCount flops for addmm.rb   
self_shaper\   r]   r,   r.   s   &&&&,r    
addmm_floprh   T   s     7$$r(   c                $    V ^8  d   QhR\         /# r5   rV   )r9   s   "r    r:   r:   Z   s      C r(   c                    V w  rEpVw  rxp	WG8w  d   \        RV RV 24      hWh8w  d   \        RV RV 24      hWE,          V	,          ^,          V,          p
V
# )z"Count flops for the bmm operation.z0bmm: batch dimensions must match (b == b2), got rY   z0bmm: inner dimensions must match (k == k2), got rZ   )r\   r]   r,   r.   br^   r_   b2r`   ra   flops   &&&,       r    bmm_floprn   Y   ss    
 GA!IBAwOPQsRWXZW[\]]wOPQsRWXZW[\]]519q=1DKr(   c                $    V ^8  d   QhR\         /# r5   rV   )r9   s   "r    r:   r:   i   s     & &C &r(   c                    \        W4      # )z&Count flops for the baddbmm operation.)rn   rf   s   &&&&,r    baddbmm_floprq   h   s    
 G%%r(   c                $    V ^8  d   QhR\         /# r5   rV   )r9   s   "r    r:   r:   p   s     % % 	%r(   c	                    \        W4      # )zCount flops for _scaled_mm.re   )
r\   r]   scale_a_shapescale_b_shape
bias_shapescale_result_shape	out_dtypeuse_fast_accumr,   r.   s
   &&&&&&&&&,r    _scaled_mm_floprz   o   s     7$$r(   c          
          V ^8  d   QhR\         \        ,          R\         \        ,          R\         \        ,          R\        R\        /# )r6   x_shapew_shaper,   
transposedr7   )listrW   bool)r9   s   "r    r:   r:      sF     $ $#Y$#Y$ Cy$ 	$
 	$r(   c                    V ^ ,          pV'       d   T MTR,          pVvrgp \        V4      \        V4      ,          V,          V,          V,          ^,          p	V	# )a  Count flops for convolution.

Note only multiplication is
counted. Computation for bias are ignored.
Flops for a transposed convolution are calculated as
flops = (x_shape[2:] * prod(w_shape) * batch_size).
Args:
    x_shape (list(int)): The input shape before convolution.
    w_shape (list(int)): The filter shape.
    out_shape (list(int)): The output shape after convolution.
    transposed (bool): is the convolution transposed
Returns:
    int: the number of flops
r6   NNr   )
r|   r}   r,   r~   
batch_size
conv_shapec_outc_infilter_sizerm   s
   &&&&      r    conv_flop_countr      sY    ( J''Y;J 'E+ 
d;//*<uDtKaODKr(   c                $    V ^8  d   QhR\         /# r5   rV   )r9   s   "r    r:   r:      s     O Oux Or(   c                   \        WWvR7      # )zCount flops for convolution.r~   )r   )
r|   r}   _bias_stride_padding	_dilationr~   r,   r-   r.   s
   &&&&&&&$*,r    	conv_flopr      s     7YNNr(   c                $    V ^8  d   QhR\         /# r5   rV   )r9   s   "r    r:   r:      s     e e er(   c                 z   R  p^ p V
^ ,          '       d+   \        V^ ,          4      pV\        WW'       * 4      ,          pV
^,          '       dm   \        V^,          4      pV'       d+   V\        V! V 4      V! V4      V! V4      RR7      ,          pV# V\        V! V4      V! V 4      V! V4      RR7      ,          pV# )c                 T    V ^,          V ^ ,          .\        V R,          4      ,           # )   r   )r   )r%   s   &r    tconv_backward_flop.<locals>.t   s"    a%(#d59o55r(   Fr   )r'   r   )grad_out_shaper|   r}   r   r   r   r   r~   _output_padding_groupsoutput_maskr,   r   
flop_countgrad_input_shapegrad_weight_shapes   &&&&&&&&&&&&    r    conv_backward_flopr      s    6JDL 1~~$Yq\2on?OQ_``
1~~%il3/!N*;QwZK\I]joppJ
  /!G*a6GK\I]joppJr(   c                l   V w  r4rVVw  rxrVw  rrY7u;8X  d   V8X  d   M MW8X  d   Wj8X  d   W8X  g   \        RV  RV RV 24      hWH8  g   WH,          ^ 8w  d   \        RV RV R24      h^ pV\        W4,          WV3W4,          Wi34      ,          pV\        W4,          WY3W4,          W34      ,          pV# )z
Count flops for self-attention.

Supports GQA (grouped-query attention) where key/value have fewer heads
than the query. The kernel broadcasts KV heads to match query heads.
z<sdpa_flop_count: query/key/value shapes are incompatible: q=z, k=z, v=zsdpa_flop_count: query heads ()) must be a multiple of key/value heads ()r[   rn   )query_shape	key_shapevalue_shaperk   h_qs_qd_q_b2h_kvs_k_d2_b3_h3_s3d_vtotal_flopss   &&&             r    sdpa_flop_countr     s     #AC#Cs$CcOO
szT)D?
 	
 zSZ1_,SE 2  $vQ(
 	
 K8QWc/!'31DEEK8QWc/!'31DEEKr(   c                $    V ^8  d   QhR\         /# r5   rV   )r9   s   "r    r:   r:   8  s     @ @WZ @r(   c                   \        WV4      # )Count flops for self-attention.r   )r   r   r   r,   r-   r.   s   &&&$*,r    	sdpa_flopr   5  s     ;;??r(   c                    ^ RI Hp ^ RIHp \	        WV34      '       g:   V P
                  P                  R8w  d   V P                  4       P                  4       # V.V P                  ^ 4      ^,
          ,          # )z
If the offsets tensor is fake, then we don't know the actual lengths.
In that case, we can just assume the worst case; each batch has max length.
)
FakeTensor)FunctionalTensormeta)
torch._subclasses.fake_tensorr   #torch._subclasses.functional_tensorr   r#   devicerH   difftolistsize)offsetsmax_lenr   r   s   &&  r    _offsets_to_lengthsr   >  s\    
 9Dg,<=>>7>>CVCVZ`C`||~$$&&9Q!+,,r(   grad_outc          	          V ^8  d   QhR\         \        \        \        R3,          \        \        R3,          \        \        R3,          \        \        R3,          R,          3,          ,          /# r6   r7   .Nr   tuplerW   )r9   s   "r    r:   r:   J  sS     1` 1` eE#s(OU38_eCHouSRUXY]G]]^_1`r(   c              #    "   VEe8   \        VP                  4      ^8w  d   \        R4      h\        VP                  4      ^8w  d   \        R4      hVe'   VP                  V P                  8w  d   \        R4      hV P                  w  rp
VP                  w  rpVP                  w  rpVf   \        R4      hVf   \        R4      hVP                  VP                  8w  d   \        R4      h\        WF4      p\        WW4      p\	        VVRR	7       F(  w  pp^V	VV
3p^VVV3p^VVV3pVe   TMRpVVVV3x  K*  	  R# V P                  VP                  VP                  Ve   VP                  MR3x  R# 5i)
a'  
Given inputs to a flash_attention_(forward|backward) kernel, this will handle behavior for
NestedTensor inputs by effectively unbinding the NestedTensor and yielding the shapes for
each batch element.

In the case that this isn't a NestedTensor kernel, then it just yields the original shapes.
Nz7sdpa_flop_count: expected key.shape to be 3-dimensionalz9sdpa_flop_count: expected value.shape to be 3-dimensionalzDsdpa_flop_count: grad_out.shape must match query.shape when providedz+sdpa_flop_count: cum_seq_q must not be Nonez+sdpa_flop_count: cum_seq_k must not be NonezAsdpa_flop_count: cum_seq_q and cum_seq_k must have the same shapeTstrictlenr%   r[   r   zip)querykeyvaluer   	cum_seq_q	cum_seq_kmax_qmax_k_r   r   h_kd_kh_vr   seq_q_lengthsseq_k_lengths	seq_q_len	seq_k_lennew_query_shapenew_key_shapenew_value_shapenew_grad_out_shapes   $$$$$$$$               r    %_unpack_flash_attention_nested_shapesr   J  sp    $  syy>Q !Z[[u{{q  !\]]HNNekk$A !ghhkkiikk !NOO !NOO??ioo- !dee+I=+I=&)-t&T"Y	 #y#6OY4M #y#6O4<4Hd!=/CUUU 'U 	
++syy%++AUx~~[_
__s   E5E7c          	          V ^8  d   QhR\         \        \        \        R3,          \        \        R3,          \        \        R3,          \        \        R3,          R,          3,          ,          /# r   r   )r9   s   "r    r:   r:   ~  sS     4` 4` eE#s(OU38_eCHouSRUXY]G]]^_4`r(   c              #    "   VEe;   \        VP                  4      ^8w  d   \        R4      h\        VP                  4      ^8w  d   \        R4      hVe'   VP                  V P                  8w  d   \        R4      hV P                  w   rp
VP                  w   rpVP                  w   rpVf   \        R4      hVf   \        R4      hVP                  VP                  8w  d   \        R4      h\        WF4      p\        WW4      p\	        VVRR	7       F(  w  pp^V	VV
3p^VVV3p^VVV3pVe   TMRpVVVV3x  K*  	  R# V P                  VP                  VP                  Ve   VP                  MR3x  R# 5i)
a+  
Given inputs to a efficient_attention_(forward|backward) kernel, this will handle behavior for
NestedTensor inputs by effectively unbinding the NestedTensor and yielding the shapes for
each batch element.

In the case that this isn't a NestedTensor kernel, then it just yields the original shapes.
NzQ_unpack_efficient_attention_nested_shapes: expected key.shape to be 4-dimensionalzS_unpack_efficient_attention_nested_shapes: expected value.shape to be 4-dimensionalz^_unpack_efficient_attention_nested_shapes: grad_out.shape must match query.shape when providedzH_unpack_efficient_attention_nested_shapes: cu_seqlens_q must not be NonezH_unpack_efficient_attention_nested_shapes: cu_seqlens_k must not be Noneza_unpack_efficient_attention_nested_shapes: cu_seqlens_q and cu_seqlens_k must have the same shapeTr   r   )r   r   r   r   cu_seqlens_qcu_seqlens_kmax_seqlen_qmax_seqlen_kr   r   r   r   r   r   r   	seqlens_q	seqlens_klen_qlen_kr   r   r   r   s   $$$$$$$$               r    )_unpack_efficient_attention_nested_shapesr   ~  s    $  syy>Q !tuuu{{q  !vwwHNNekk$A   "B  C  C131313 !kll !kll!3!33  "Z [ ['C	'C		9TBLE5 #uc2OUC0M #uc2O4<4Hd!=/CUUU C 	
++syy%++AUx~~[_
__s   E8E:)rQ   c                $    V ^8  d   QhR\         /# r5   rV   )r9   s   "r    r:   r:            	r(   c          
     J    \        V VVVVVVR7      p
\        R V
 4       4      # )r   r   r   r   r   r   r   r   c              3   @   "   T F  w  rr4\        WV4      x  K  	  R # 5ir   r   r   r   r   r   r   s   &    r    r!   0_flash_attention_forward_flop.<locals>.<genexpr>  &      6;2KK 	<<6;   r   sum)r   r   r   r   r   r   r   r,   r-   r.   sizess   &&&&&&&$*, r    _flash_attention_forward_flopr     s?    " 2E  6;  r(   c                $    V ^8  d   QhR\         /# r5   rV   )r9   s   "r    r:   r:     r   r(   c           
     J    \        V VVVVVVR7      p
\        R V
 4       4      # )r   )r   r   r   r   r   r   r   c              3   @   "   T F  w  rr4\        WV4      x  K  	  R # 5ir   r   r   s   &    r    r!   4_efficient_attention_forward_flop.<locals>.<genexpr>  r   r   r   r   )r   r   r   biasr   r   r   r   r-   r.   r   s   &&&&&&&&*, r    !_efficient_attention_forward_flopr    s?    " 6!!!!E  6;  r(   c                    Vw  rErgVw  rrVw  rrV w  ppppYHu;8X  d   Tu;8X  d   V8X  d   M MW8X  d   VV8X  g   \        R 4      hWY8  g   WY,          ^ 8w  d   \        RV RV	 R24      hW{8X  d   VV8X  d   W8X  d   VV8X  g   \        R4      h^ pV\        WE,          Wg3WE,          Wz34      ,          pV\        WE,          Wo3WE,          W34      ,          pV\        WE,          W3WE,          Wo34      ,          pV\        WE,          Wj3WE,          W34      ,          pV\        WE,          Wv3WE,          Wj34      ,          pV# )z<sdpa_backward_flop_count: batch/heads mismatch among tensorsz'sdpa_backward_flop_count: query heads (r   r   zJsdpa_backward_flop_count: grad_out/value/key/query shapes are incompatibler   )r   r   r   r   rk   r   r   r   r   r   r   r   r   r   r   r   _b4_h4_s4_d4r   s   &&&&                 r    sdpa_backward_flop_countr    sa   "AC#Cs$Cc'Cc3""s"t{sczJ
 	
 zSZ1_5cU ;  $vQ(
 	
 J3#:#*X
 	
 K 8QWc/!'31DEEK 8QWc/!'31DEEK8QWc/!'31DEEK 8QWc/!'31DEEK8QWc/!'31DEEKr(   c                $    V ^8  d   QhR\         /# r5   rV   )r9   s   "r    r:   r:     s     Y Yps Yr(   c                   \        WW#4      # )z(Count flops for self-attention backward.r  )r   r   r   r   r,   r-   r.   s   &&&&$*,r    sdpa_backward_flopr    s    
 $NXXr(   c                $    V ^8  d   QhR\         /# r5   rV   )r9   s   "r    r:   r:   "         	r(   c
                 L    \        VVVV VVVV	R 7      p\        R V 4       4      # )r   r   r   r   r   r   r   r   c              3   @   "   T F  w  rr4\        WAW#4      x  K  	  R # 5ir   r  r   r   r   r   r   s   &    r    r!   1_flash_attention_backward_flop.<locals>.<genexpr><  &      CI?KK 	!iUUCIr   r   )r   r   r   r   out	logsumexpr   r   r   r   r-   r.   shapess   &&&&&&&&&&*, r    _flash_attention_backward_flopr  !  sB    " 3	F  CI  r(   c                $    V ^8  d   QhR\         /# r5   rV   )r9   s   "r    r:   r:   C  r  r(   c
                 L    \        VVVV VVVV	R 7      p\        R V 4       4      # ))r   r   r   r   r   r   r   r   c              3   @   "   T F  w  rr4\        WAW#4      x  K  	  R # 5ir   r  r  s   &    r    r!   5_efficient_attention_backward_flop.<locals>.<genexpr>]  r  r   r   )r   r   r   r   r  r  r   r   r   r   r-   r.   r  s   &&&&&&&&&&*, r    "_efficient_attention_backward_flopr  B  sB    " 7!!!!	F  CI  r(   r*   c                $    V ^8  d   QhR\         /# r5   rV   )r9   s   "r    r:   r:   c  s       	r(   c          
     V    \        T TTTVe   TMTVVR7      p
\        R V
 4       4      # )z$Count flops for varlen_attn forward.r   c              3   @   "   T F  w  rr4\        WV4      x  K  	  R # 5ir   r   r   s   &    r    r!   ,_varlen_attn_forward_flop.<locals>.<genexpr>y  r   r   r   )r   r   r   cu_seq_qcu_seq_kr   r   r*   r-   r.   r   s   &&&&&&&$*, r    _varlen_attn_forward_flopr$  c  sF     2&2(E  6;  r(   c                $    V ^8  d   QhR\         /# r5   rV   )r9   s   "r    r:   r:     s       	r(   c          	         \        WW4WVV4      # )z(Count flops for varlen_attn_out forward.)r$  )r  r   r   r   r"  r#  r   r   r*   r-   r.   s   &&&&&&&&$*,r    _varlen_attn_out_flopr'    s     %EXe r(   c                $    V ^8  d   QhR\         /# r5   rV   )r9   s   "r    r:   r:     s       	r(   c
               L    \        VVVV VVVV	R7      p\        R V 4       4      # )z%Count flops for varlen_attn backward.r  c              3   @   "   T F  w  rr4\        WAW#4      x  K  	  R # 5ir   r  r  s   &    r    r!   -_varlen_attn_backward_flop.<locals>.<genexpr>  s&      CH?KK 	!iUUCHr   r   )r   r   r   r   r  lser"  r#  r   r   r*   r-   r.   r   s   &&&&&&&&&&$*, r    _varlen_attn_backward_flopr-    sB      2	E  CH  r(   c                 8    \        V \        4      '       g   V 3# V # r   )r#   r   )xs   &r    normalize_tupler0    s    atHr(   c                     \        ^ \        \        \        4      ^,
          \        \	        V 4      4      ^,
          ^,          4      4      p\        V,          # )r   )maxminr   suffixesstr)numberindexs   & r    get_suffix_strr8    s=     3s8}q(3s6{+;a+?A*EFGEE?r(   c                 x    \         P                  V4      pV R V,          ,          R pV\         V,          ,           # )i  z.3f)r4  r7  )r6  suffixr7  r   s   &&  r    convert_num_with_suffixr;    s2    NN6"E%c*E8E?""r(   c                $    V ^8  d   QhR\         /# r5   )r5  )r9   s   "r    r:   r:     s        #  r(   c                 (    V^ 8X  d   R# W,          R # )r   0%z.2%rA   )numdenoms   &&r    convert_to_percent_strrA    s    zk#r(   c                 0   a  \        S 4      V 3R  l4       pV# )c                 @   < \        V 4      w  rS! V!  p\        W24      # r   )r   r   )r-   	flat_argsspecr  r/   s   &   r    r0   )_pytreeify_preserve_structure.<locals>.nf  s#    &t,	mc((r(   r   r2   s   f r    _pytreeify_preserve_structurerG    s     
1X) )
 Ir(   c                      a a ] tR tRt oRtRV3R lV 3R llltV3R lR ltV3R lR ltRR	 ltR
 t	R t
R tRtVtV ;t# )r   i  a  
``FlopCounterMode`` is a context manager that counts the number of flops within its context.

It does this using a ``TorchDispatchMode``.

It also supports hierarchical output by passing a module (or list of
modules) to FlopCounterMode on construction. If you do not need hierarchical
output, you do not need to use it with a module.

Example usage

.. code-block:: python

    mod = ...
    with FlopCounterMode(mod) as flop_counter:
        mod.sum().backward()

c          
         < V ^8  d   QhRS[ P                  P                  S[S[ P                  P                  ,          ,          R,          RS[RS[RS[S[S[3,          R,          RR/# )r6   modsNdepthdisplaycustom_mappingr7   )r   nnModuler   rW   r   dictr   )r9   __classdict__s   "r    r:   FlopCounterMode.__annotate__  sj     + +((//D$99D@+ + 	+
 !cNT1+
 >B+r(   c                z  < \         SV `  4        \        R  4      V n        W n        W0n        RV n        Vf   / pVe   \        P                  ! R^R7       / \        CVP                  4        UUu/ uF&  w  rVT\        VRR4      '       d   TM
\        V4      bK(  	  uppCV n	        \        4       V n        R# u uppi )c                       \        \        4      # r   )r   rW   rA   r(   r    <lambda>*FlopCounterMode.__init__.<locals>.<lambda>  s
    +VYJZr(   Nz<mods argument is not needed anymore, you can stop passing it)
stacklevel_get_rawF)super__init__r   flop_countsrK  rL  modewarningswarnrI   itemsr   r3   r   mod_tracker)selfrJ  rK  rL  rM  r_   v	__class__s   &&&&&  r    rZ  FlopCounterMode.__init__  s     	6ABZ6[
-1	!NMMXefg

WeWkWkWmnWmtqqwq*e44!-:JJWmn
 )? os   0,B7c                    < V ^8  d   QhRS[ /# r5   rV   )r9   rQ  s   "r    r:   rR    s     8 8 8r(   c                V    \        V P                  R ,          P                  4       4      # )Global)r   r[  valuesra  s   &r    get_total_flopsFlopCounterMode.get_total_flops  s!    4##H-44677r(   c                L   < V ^8  d   QhRS[ S[S[ S[S[3,          3,          /# r5   )rP  r5  r   rW   )r9   rQ  s   "r    r:   rR    s)     
A 
Ac4S>&9!: 
Ar(   c                ~    V P                   P                  4        UUu/ uF  w  rV\        V4      bK  	  upp# u uppi )zReturn the flop counts as a dictionary of dictionaries.

The outer
dictionary is keyed by module name, and the inner dictionary is keyed by
operation name.

Returns:
    Dict[str, Dict[Any, int]]: The flop counts as a dictionary.
)r[  r_  rP  )ra  r_   rb  s   &  r    get_flop_countsFlopCounterMode.get_flop_counts  s7     (,'7'7'='='?@'?tq47
'?@@@s   9c                d  a a
aa Vf   S P                   pVf   Rp^ R IpRVn        . R	Op. pS P                  4       o
\	        S
4      oRoV
VVV 3R lp\        S P                  P                  4       4       FL  pVR8X  d   K  VP                  R4      ^,           pWq8  d   K,  V! Wg^,
          4      pVP                  V4       KN  	  RS P                  9   d5   S'       g-   V F  p	RV	^ ,          ,           V	^ &   K  	  V! R^ 4      V,           p\        V4      ^ 8X  d   . R
O.pVP                  WCRR7      # )Ni?B TFc           	        < \        S
P                  V ,          P                  4       4      pS	VS8  ,          o	R V,          p. pVP                  W0,           \	        VS4      \        VS4      .4       S
P                  V ,          P                  4        FD  w  rVVP                  VR,           \        V4      ,           \	        VS4      \        VS4      .4       KF  	  V# ) z - )r   r[  rh  appendr;  rA  r_  r5  )mod_namerK  r   paddingrh  r_   rb  global_flopsglobal_suffixis_global_subsumedra  s   &&     r    process_mod.FlopCounterMode.get_table.<locals>.process_mod8  s     d..x8??ABK+"==EkGFMM"']C&{LA 
 ((288:eOc!f,+A}=*1l;  ; Mr(   rg  .rr  )headerscolalign)rO  FLOPz% Total)rg  0r>  )leftrightr  )rK  tabulatePRESERVE_WHITESPACErj  r8  sortedr[  keyscountextendr   )ra  rK  r  headerrh  ry  mod	mod_depth
cur_valuesr   rv  rw  rx  s   f&        @@@r    	get_tableFlopCounterMode.get_table(  s%   =JJE=E 	'+$.++-&|4"	 	, $**//12Ch		#*I $Sa-8JMM*% 3 t'''0Bq>a   !1-6Fv;!+,F  B\ ]]r(   c                    V P                   P                  4        V P                  P                  4        \	        V 4      V n        V P
                  P                  4        V # r   )r[  clearr`  	__enter___FlopCounterModer\  ri  s   &r    r  FlopCounterMode.__enter__g  sG     ""$$T*			r(   c                   V P                   f   \        R4      hV P                   P                  ! V!  pR V n         V P                  P                  4        V P                  '       d%   \        V P                  V P                  4      4       V# )Nz<Internal error: FlopCounter.__exit__ called but mode is None)r\  r[   __exit__r`  rL  printr  rK  )ra  r-   rk   s   &* r    r  FlopCounterMode.__exit__n  sh    99 !_``II%	!!#<<<$..,-r(   c                    WP                   9   dl   V P                   V,          pV! V/ VBR V/B p\        V P                  P                  4       F)  pV P                  V,          V;;,          V,          uu&   K+  	  V# )r*   )rI   setr`  parentsr[  )ra  func_packetr  r-   r.   flop_count_funcr   pars   &&&&&   r    _count_flopsFlopCounterMode._count_flopsx  sm    ,,,"00=O($F&F#FJ4++334  %k2j@2 5
r(   )rK  rL  r[  rI   r`  r\  )Nr6   TNr   )__name__
__module____qualname____firstlineno____doc__rZ  rj  rn  r  r  r  r  __static_attributes____classdictcell____classcell__)rc  rQ  s   @@r    r   r     sE     &+ +*8 8
A 
A<^~ r(   c                   L   a  ] tR tRt o RtV 3R lR ltR tR tR
R ltR	t	V t
R# )r  i  Tc                $   < V ^8  d   QhRS[ RR/# )r6   counterr7   N)r   )r9   rQ  s   "r    r:   _FlopCounterMode.__annotate__  s       D r(   c                    Wn         R # r   r  )ra  r  s   &&r    rZ  _FlopCounterMode.__init__  s    r(   c                   ^ RI pVP                  V P                  P                  4      pV ;_uu_ 4        V! V!  pRRR4       VP                  V P                  P                  4      pW@P                  n        XV3#   + '       g   i     LI; i)a]  Execute a branch function and capture its FLOP counts without
affecting self.counter.flop_counts

Args:
    branch_fn: The branch function to execute
    operands: Arguments to pass to the branch function

Returns:
    Tuple of (result, flop_counts) where result is the branch output
    and flop_counts is a copy of the FLOP counts after execution
N)copyr  r[  )ra  	branch_fnoperandsr  checkpointed_flop_countsresultr[  s   &&&    r    $_execute_with_isolated_flop_counting5_FlopCounterMode._execute_with_isolated_flop_counting  si     	#'99T\\-E-E#F T)F ii 8 89#; {""	 Ts   A<<B	c                   V\         P                  P                  P                  \         P                  P                  P                  09   pV'       dk   ^ RIHp ^ RIHp V! VR,          4      p\        W4      '       g"   \        VR4      '       d   VP                  pK1   V P                  P                  VRW44      # V\         P                  P                  P                  J Edc   Vw  rrV P                  W4      w  rV\         J d   \         # V P                  W4      w  ppV\         J d   \         # \#        VP%                  4       4      \#        VP%                  4       4      ,          p/ pV F  pVV,          pVV,          p/ p\#        VP%                  4       4      \#        VP%                  4       4      ,          pV F6  pVP'                  V^ 4      pVP'                  V^ 4      p\)        VV4      VV&   K8  	  VVV&   K  	  VP+                  4        F2  w  ppV P                  P,                  V,          P/                  V4       K4  	  V# \         # )r   )
get_kernelr   
kernel_idxfnN)r   opshigher_ordertriton_kernel_wrapper_mutation triton_kernel_wrapper_functional*torch._higher_order_ops.triton_kernel_wrapr  triton.runtime.jitr   r#   hasattrr  r  r  condr  NotImplementedr  r  getr2  r_  r[  update)ra  functypesr-   r.   	is_tritonr  r   kernel_namepredtrue_branchfalse_branchr  true_outtrue_flop_counts	false_outfalse_flop_countsall_mod_keysmerged_flop_counts	outer_keytrue_func_countsfalse_func_countsmerged_func_countsall_func_keysfunc_keytrue_val	false_val
inner_dicts   &&&&&                       r    _handle_higher_order_ops)_FlopCounterMode._handle_higher_order_ops  s   UYY33RR"YY33TTV V	M6$VL%9:K ::;--"-..K<<,,[$MMUYY++000
 9=5D|)-)R)R*&H >)%%+/+T+T,(I( N*%% /4467#>O>T>T>V:WWL!#)	#3I#> $5i$@!%'" #$4$9$9$; <sCTCYCYC[?\ \ -H/33Ha@H 1 5 5h BI36x3K&x0 !.
 1C"9- * *<)A)A)C%	:((3:::F *D
 O!!r(   Nc                   V'       d   TM/ pV\         P                  P                  P                  P                  \         P                  P                  P
                  P                  \         P                  P                  P
                  P                  \         P                  P                  P                  P                  \         P                  P                  P                  P                  \         P                  P                  P                  P                  \         P                  P                  P                  P                  \         P                  P                  P                  P                  \         P                  P                  P                  P                  \         P                  P                  P                  P                  \         P                  P                  P                  P                  \         P                  P                  P                  P                  \         P                  P                  P                   P                  \         P                  P                  P"                  P                  \         P                  P$                  P&                  P                  09   d   \(        # \+        V\         P,                  P.                  4      '       d   V P1                  WW44      # WP2                  P4                  9  do   V\         P                  P$                  P6                  P                  Jd=   T ;_uu_ 4        VP8                  ! V/ VB pV\(        Jd   VuuR R R 4       #  R R R 4       V! V/ VB pV P2                  P;                  VP<                  WcV4      #   + '       g   i     L?; ir   )r   r  atensym_is_contiguousdefaultis_contiguousmemory_formatis_strides_like_formatis_non_overlapping_and_denser   sym_sizestride
sym_stridestorage_offsetsym_storage_offsetnumel	sym_numeldimprimlayoutr  r#   rD   HigherOrderOperatorr  r  rI   r   	decomposer  _overloadpacket)ra  r  r  r-   r.   rr  s   &&&&&  r    __torch_dispatch__#_FlopCounterMode.__torch_dispatch__  s;   !r EIINN44<<IINN0088IINN00>>IINN99AAIINN??GGIINN''//IINN++33IINN))11IINN--55IINN1199IINN55==IINN((00IINN,,44IINN&&..IINN))113 3  "!dEJJ::;;00dKK ||111d%))..BWBWB_B_6_NND3F3N* *  D#F#||(()=)=s&QQ s   N00O 	r  )rA   N)r  r  r  r  supports_higher_order_operatorsrZ  r  r  r  r  r  )rQ  s   @r    r  r    s,     &*# #(;"z"R "Rr(   r  c                b    V ^8  d   Qh/ ^ \         9   d   \        \        \        3,          ;R&   # )r6   rI   )__conditional_annotations__rP  r   )r9   s   "r    r:   r:      s$      L # "tCH~ "M r(   )cudahipxpu)Fr   )NNNFN) KMBT)br  r  r   loggingr   torch.utils._pytreer   r   r   module_trackerr   typingr   r	   collections.abcr
   r   typing_extensionsr   collectionsr   torch.utils._python_dispatchr   mathr   	functoolsr   r]  __all__r   r   	getLoggerr  logr  r   rF   ImportErroranywarningr  r  r'   rI   r3   r   mmrb   addmmrh   bmmrn   baddbmmrq   
_scaled_mmrz   r   convolution_convolutioncudnn_convolution_slow_conv2d_forwardconvolution_overrideabler   convolution_backwardr   r   '_scaled_dot_product_efficient_attention#_scaled_dot_product_flash_attention#_scaled_dot_product_cudnn_attentionr   r   r   r   _flash_attention_forwardr   _efficient_attention_forwardr  r  0_scaled_dot_product_efficient_attention_backward,_scaled_dot_product_flash_attention_backward,_scaled_dot_product_cudnn_attention_backwardr  _flash_attention_backwardr  _efficient_attention_backwardr  r$  r'  r-  r0  r4  r8  r;  rA  rG  r   r  r:   )r  s   @r    <module>r%     s       F F )  $ $ ' # :   5
6T]t_!> yy~~
 !# ". tww	t 	  	 tzz"% #% txx  ! t||$& %& t'% (% $L (())..1155	7 8
Obf O8
O t001e 2eN8 DD@@@@B C@D @C@	-1`
 1`h4`
 4`n t44dC  D> t88$G H>"J MMIIIIK LY]a YLY t55tD E@ t994H I@ 8 & @GGWJJ
 	HHh 	LL,	
 	OO_ 	i 	y 	I 	!!9 	y 	1 	00) 	,,i 	,,i 	99;M  	557I!" 	557I#$ 	!!#@%%'H""$B&&(J+0 $# 
N N`yR( yRK  

]F\
]sss
]F\
]]]VWLs$   P Q%Q%3Q%Q%$Q%