ó
    "EñiòŠ  ã                   ó,	  • % S SK Jr  S SKrS SKrS SKJrJrJr  SSKJ	r	  S SK
JrJr  S SKJr  S SKJr  S S	KJr  S S
KJr  S SKJr  S SKJr  S SKJr  S SKrSS/r\" S5      r\" S5      r\R<                  " \5      r  S SK!J"r#  \RN                  RP                  r(S r)0 r*\+\\4   \,S'   S r-S@S\\\\4   /\\\4   4   4S jjr.\." \(R^                  5      SS.S\04S jj5       r1\." \(Rd                  5      SAS\04S jj5       r3\." \(Rh                  5      SAS\04S jj5       r5\." \(Rl                  5      SAS\04S jj5       r7\." \(Rp                  5           SBS\04S  jj5       r9 S@S!\:\0   S"\:\0   S#\:\0   S$\;S\04
S% jjr<\." \(Rz                  \(R|                  \(R~                  \(R€                  \(R‚                  /5      SS.S\04S& jj5       rB\." \(R†                  5      S\04S' j5       rDS( rE\." \(RŒ                  \(RŽ                  \(R�                  /5      SS.S\04S) jj5       rIS* rJSS+.S\\K\K\0S,4   \K\0S,4   \K\0S,4   \K\0S,4   S-  4      4S- jjrLSS+.S\\K\K\0S,4   \K\0S,4   \K\0S,4   \K\0S,4   S-  4      4S. jjrM\." \(Rœ                  S/S09SS.S\04S1 jj5       rO\." \(R                   S/S09S\04S2 j5       rQS3 rR\." \(R¦                  \(R¨                  \(Rª                  /5      SS.S\04S4 jj5       rV\." \(R®                  S/S09S\04S5 j5       rX\." \(R²                  S/S09S\04S6 j5       rZ0 \(R^                  \1_\(Rd                  \3_\(Rh                  \5_\(Rl                  \7_\(Rp                  \9_\(Rz                  \B_\(R|                  \B_\(R~                  \B_\(R‚                  \B_\(R€                  \B_\(R†                  \D_\(RŒ                  \I_\(RŽ                  \I_\(R�                  \I_\(R¦                  \V_\(R¨                  \V_\(Rª                  \V_\(Rœ                  \O\(R                   \Q\(R®                  \X\(R²                  \Z0Er*S7 r[/ S8Qr\S9 r]S: r^S\_4S; jr`S< ra " S= S5      rb " S> S?\5      rcg! \$ a+    \%" S S 5       5      (       a  \ RM                  S5        \r# GNf = f)Cé    )ÚNoneTypeN)Útree_mapÚtree_flattenÚtree_unflattené   )ÚModuleTracker)ÚAnyÚTypeVar)ÚCallable)ÚIterator)Ú	ParamSpec)Údefaultdict)ÚTorchDispatchMode©Úprod©ÚwrapsÚFlopCounterModeÚregister_flop_formulaÚ_TÚ_P©ÚJITFunctionc              #   ó\   #   • U  H"  n[        [        R                  US 5      S Lv •  M$     g 7f©N)ÚgetattrÚtorchÚversion)Ú.0Úattrs     ÚU/home/mande/repo/quber/.venv/lib/python3.13/site-packages/torch/utils/flop_counter.pyÚ	<genexpr>r"      s$   é € Ð
]ÒF\¸dŒ7”5—=‘= $¨Ó-°TÕ9ÒF\ùs   ‚*,)ÚcudaÚhipÚxpuz@triton not found; flop counting will not work for triton kernelsc                 ó\   • [        U [        R                  5      (       a  U R                  $ U $ r   )Ú
isinstancer   ÚTensorÚshape)Úis    r!   Ú	get_shaper+   #   s!   € Ü�!”U—\‘\×"Ñ"Ø�w‰wˆØ€Hó    Úflop_registryc                 ó8   ^ • [        T 5      S S.U 4S jj5       nU$ )N)Úout_valc                 óB   >• [        [        XU 45      u  pnT" USU0UD6$ )NÚ	out_shape)r   r+   )r/   ÚargsÚkwargsr1   Úfs       €r!   ÚnfÚshape_wrapper.<locals>.nf+   s.   ø€ ä"*¬9°tÀWÐ6MÓ"NÑˆ�iÙ�$Ð6 )Ð6¨vÑ6Ð6r,   r   ©r4   r5   s   ` r!   Úshape_wrapperr8   *   s#   ø€ Ü
ˆ1ƒXØ÷ 7ó ð7ð €Ir,   Úreturnc                 óh   ^ ^• S[         [        [        4   S[         [        [        4   4UU 4S jjnU$ )NÚflop_formular9   c                 ó’   >^ • T(       d  [        T 5      m SU 4S jjn[        R                  R                  R	                  UT5        T $ )Nc                 óÚ   >• [        U [        R                  R                  [        45      (       d  [        SU  S[        U 5       35      eU [        ;   a  [        SU  35      eT[        U '   g )Nz|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 )	r'   r   Ú_opsÚOpOverloadPacketÚ_JITFunctionÚ
ValueErrorÚtyper-   ÚRuntimeError)Útargetr;   s    €r!   ÚregisterÚ=register_flop_formula.<locals>.register_fun.<locals>.register7   sp   ø€ Ü˜v¬¯
©
×(CÑ(CÄ\Ð'R×SÑSÜ ðà#˜HÐ$6´t¸F³|°nðFóGð Gð œÓ&Ü"Ð%AÀ&ÀÐ#JÓKÐKØ$0ŒM˜&Ò!r,   )r9   N)r8   r   ÚutilsÚ_pytreeÚ	tree_map_)r;   rE   Úget_rawÚtargetss   ` €€r!   Úregister_funÚ+register_flop_formula.<locals>.register_fun3   s7   ù€ ÞÜ(¨Ó6ˆL÷	1ô 	�‰×Ñ×%Ñ% h°Ô8àÐr,   )r   r   r   )rK   rJ   rL   s   `` r!   r   r   1   s5   ù€ ð¤8¬B´¨FÑ#3ð ¼ÄÄRÀÑ8H÷ ð ð& Ðr,   )r1   c                óR   • U u  pVUu  pxXg:w  a  [        SU SU 35      eXX-  S-  U-  $ )zCount flops for matmul.z3matmul: inner dimensions must match (k == k2), got ú and é   ©ÚAssertionError)	Úa_shapeÚb_shaper1   r2   r3   ÚmÚkÚk2Úns	            r!   Úmm_floprY   H   sE   € ð
 �D€AØ�E€BØƒwÜÐRÐSTÐRUÐUZÐ[]ÐZ^Ð_Ó`Ð`à‰5�1‰9�q‰=Ðr,   c                 ó   • [        X5      $ )zCount flops for addmm.©rY   ©Ú
self_shaperS   rT   r1   r3   s        r!   Ú
addmm_flopr^   T   s   € ô �7Ó$Ð$r,   c                 óŒ   • U u  pEnUu  pxn	XG:w  a  [        SU SU 35      eXh:w  a  [        SU SU 35      eXE-  U	-  S-  U-  n
U
$ )z"Count flops for the bmm operation.z0bmm: batch dimensions must match (b == b2), got rO   z0bmm: inner dimensions must match (k == k2), got rP   rQ   )rS   rT   r1   r3   ÚbrU   rV   Úb2rW   rX   Úflops              r!   Úbmm_floprc   Y   ss   € ð
 �G€Aˆ!Ø�I€BˆAØƒwÜÐOÐPQÈsÐRWÐXZÐW[Ð\Ó]Ð]ØƒwÜÐOÐPQÈsÐRWÐXZÐW[Ð\Ó]Ð]à‰5�1‰9�q‰=˜1Ñ€DØ€Kr,   c                 ó   • [        X5      $ )z&Count flops for the baddbmm operation.)rc   r\   s        r!   Úbaddbmm_flopre   h   s   € ô
 �GÓ%Ð%r,   c	                 ó   • [        X5      $ )zCount flops for _scaled_mm.r[   )
rS   rT   Úscale_a_shapeÚscale_b_shapeÚ
bias_shapeÚscale_result_shapeÚ	out_dtypeÚuse_fast_accumr1   r3   s
             r!   Ú_scaled_mm_floprm   o   s   € ô �7Ó$Ð$r,   Úx_shapeÚw_shaper1   Ú
transposedc                 ó|   • U S   nU(       a  U OUSS nUtpgn [        U5      [        U5      -  U-  U-  U-  S-  n	U	$ )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
r   rP   Nr   )
rn   ro   r1   rp   Ú
batch_sizeÚ
conv_shapeÚc_outÚc_inÚfilter_sizerb   s
             r!   Úconv_flop_countrw   €   s[   € ð( ˜‘€JÞ'‘'¨Y¸¸Ð;€JØ 'Ð€E�+ðô �
Óœd ;Ó/Ñ/°*Ñ<¸uÑDÀtÑKÈaÑO€DØ€Kr,   c                ó   • [        XXvS9$ )zCount flops for convolution.©rp   )rw   )
rn   ro   Ú_biasÚ_strideÚ_paddingÚ	_dilationrp   r1   r2   r3   s
             r!   Ú	conv_flopr~   ¦   s   € ô ˜7¨YÑNÐNr,   c                 ó0  • S nSn U
S   (       a"  [        US   5      nU[        XXç(       + 5      -  nU
S   (       aY  [        US   5      nU(       a#  U[        U" U 5      U" U5      U" U5      SS9-  nU$ U[        U" U5      U" U 5      U" U5      SS9-  nU$ )Nc                 ó4   • U S   U S   /[        U SS  5      -   $ )Nr   r   rP   )Úlist)r)   s    r!   ÚtÚconv_backward_flop.<locals>.tÀ   s$   € Ø�a‘˜% ™(Ð#¤d¨5°°¨9£oÑ5Ð5r,   r   r   Fry   )r+   rw   )Úgrad_out_shapern   ro   rz   r{   r|   r}   rp   Ú_output_paddingÚ_groupsÚoutput_maskr1   r‚   Ú
flop_countÚgrad_input_shapeÚgrad_weight_shapes                   r!   Úconv_backward_flopr‹   ±   s³   € ò6à€JðDðL �1‡~Ü$ Y¨q¡\Ó2ÐØ”o nÐ?OÔQ_Ó`Ñ`ˆ
à�1‡~Ü% i°¡lÓ3ÐÞàœ/©!¨NÓ*;¹Q¸w»ZÉÐK\ÓI]ÐjoÑpÑpˆJð
 Ðð œ/©!¨G«*±a¸Ó6GÉÐK\ÓI]ÐjoÑpÑpˆJàÐr,   c                 óô   • U u  p4pVUu  pxpšUu  p¼pÞX7s=:X  a  U:X  a!  O  OXHs=:X  a  U:X  a  O  OXj:X  a
  X�:X  a  Xj:X  d  [        S5      eSnU[        X4-  XV4X4-  Xi45      -  nU[        X4-  XY4X4-  Xž45      -  nU$ )zR
Count flops for self-attention.

NB: We can assume that value_shape == key_shape
z8sdpa_flop_count: query/key/value shapes are incompatibler   ©rR   rc   )Úquery_shapeÚ	key_shapeÚvalue_shaper`   ÚhÚs_qÚd_qÚ_b2Ú_h2Ús_kÚ_d2Ú_b3Ú_h3Ú_s3Úd_vÚtotal_flopss                   r!   Úsdpa_flop_countr�     s•   € ð !�N€Aˆ#Ø"Ñ€CˆcØ$Ñ€CˆcØ�?�sŽ? !¥/¨c¦/¸»È3Ë:Ð]`Ó]gÜÐWÓXÐXØ€Kà”8˜Q™U CÐ-°±°sÐ/@ÓAÑA€Kà”8˜Q™U CÐ-°±°sÐ/@ÓAÑA€KØÐr,   c                ó   • [        XU5      $ )úCount flops for self-attention.©r�   )rŽ   r�   r�   r1   r2   r3   s         r!   Ú	sdpa_flopr¡   ,  s   € ô ˜;°;Ó?Ð?r,   c                 óÞ   • SSK Jn  SSKJn  [	        XU45      (       d8  U R
                  R                  S:w  a  U R                  5       R                  5       $ U/U R                  S5      S-
  -  $ )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.
r   )Ú
FakeTensor)ÚFunctionalTensorÚmetar   )
Útorch._subclasses.fake_tensorr£   Ú#torch._subclasses.functional_tensorr¤   r'   ÚdevicerB   ÚdiffÚtolistÚsize)ÚoffsetsÚmax_lenr£   r¤   s       r!   Ú_offsets_to_lengthsr®   5  s\   € õ
 9ÝDÜ�gÐ,<Ð=×>Ñ>À7Ç>Á>×CVÑCVÐZ`ÓC`Ø�|‰|‹~×$Ñ$Ó&Ð&Øˆ9˜Ÿ™ Q›¨!Ñ+Ñ,Ð,r,   )Úgrad_out.c              #   óÔ  #   • UGb+  [        UR                  5      S:w  a  [        S5      e[        UR                  5      S:w  a  [        S5      eUb%  UR                  U R                  :w  a  [        S5      eU R                  u  p‰n
UR                  u  p‹nUR                  u  p�nUc  [        S5      eUc  [        S5      eUR                  UR                  :w  a  [        S5      e[        XF5      n[        XW5      n[	        UUS	S
9 H'  u  nnSU	UU
4nSUUU4nSUUU4nUb  UOSnUUUU4v •  M)     gU R                  UR                  UR                  Ub  UR                  OS4v •  g7f)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.
Né   z7sdpa_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©Ústrictr   ©Úlenr)   rR   r®   Úzip)ÚqueryÚkeyÚvaluer¯   Ú	cum_seq_qÚ	cum_seq_kÚmax_qÚmax_kÚ_Úh_qr“   Ú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Ë   A  sp  é € ð$ Òô ˆs�y‰y‹>˜QÓÜ Ð!ZÓ[Ð[Üˆu�{‰{Ó˜qÓ Ü Ð!\Ó]Ð]ØÑ H§N¡N°e·k±kÓ$AÜ Ð!gÓhÐhØ—k‘k‰ˆ�Ø—i‘i‰ˆ�Ø—k‘k‰ˆ�ØÑÜ Ð!NÓOÐOØÑÜ Ð!NÓOÐOØ�?‰?˜iŸo™oÓ-Ü Ð!dÓeÐeÜ+¨IÓ=ˆÜ+¨IÓ=ˆÜ&)¨-¸ÈtÔ&TÑ"ˆY˜	Ø  # y°#Ð6ˆOØ  Y°Ð4ˆMØ  # y°#Ð6ˆOØ4<Ñ4H¡ÈdÐØ! =°/ÐCUÐUÔUñ 'Uð 	à
�+‰+�s—y‘y %§+¡+ÀÑAU¨x¯~ª~Ð[_Ð
_Ó_ùs   ‚E&E(c              #   óÚ  #   • UGb.  [        UR                  5      S:w  a  [        S5      e[        UR                  5      S:w  a  [        S5      eUb%  UR                  U R                  :w  a  [        S5      eU R                  u    p‰n
UR                  u    p‹nUR                  u    p�nUc  [        S5      eUc  [        S5      eUR                  UR                  :w  a  [        S5      e[        XF5      n[        XW5      n[	        UUS	S
9 H'  u  nnSU	UU
4nSUUU4nSUUU4nUb  UOSnUUUU4v •  M)     gU R                  UR                  UR                  Ub  UR                  OS4v •  g7f)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.
Né   zQ_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¹   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Ö   u  s‹  é € ð$ Òô ˆs�y‰y‹>˜QÓÜ Ð!tÓuÐuÜˆu�{‰{Ó˜qÓ Ü Ð!vÓwÐwØÑ H§N¡N°e·k±kÓ$AÜ ð  "Bó  Cð  CØŸ™‰ˆˆ1�3ØŸ™‰ˆˆ1�3ØŸ™‰ˆˆ1�3ØÑÜ Ð!kÓlÐlØÑÜ Ð!kÓlÐlØ×Ñ ×!3Ñ!3Ó3Ü ð "Zó [ð [ä'¨ÓCˆ	Ü'¨ÓCˆ	Ü 	¨9¸TÔB‰LˆE�5Ø  # u¨cÐ2ˆOØ  U¨CÐ0ˆMØ  # u¨cÐ2ˆOØ4<Ñ4H¡ÈdÐØ! =°/ÐCUÐUÔUñ Cð 	à
�+‰+�s—y‘y %§+¡+ÀÑAU¨x¯~ª~Ð[_Ð
_Ó_ùs   ‚E)E+T)rJ   c          
      óD   • [        U UUUUUUS9n
[        S U
 5       5      $ )rŸ   )r·   r¸   r¹   rº   r»   r¼   r½   c              3   ó@   #   • U  H  u  pp4[        XU5      v •  M     g 7fr   r    ©r   rŽ   r�   r�   r¾   s        r!   r"   Ú0_flash_attention_forward_flop.<locals>.<genexpr>Æ  ó&   é € ð â6;Ñ2ˆK Kô 	˜°×<Ð<Ú6;ùó   ‚©rË   Úsum)r·   r¸   r¹   rº   r»   r¼   r½   r1   r2   r3   Úsizess              r!   Ú_flash_attention_forward_floprà   ¬  s?   € ô" 2ØØØØØØØñ€Eô ñ á6;óó ð r,   c           
      óD   • [        U UUUUUUS9n
[        S U
 5       5      $ )rŸ   )r·   r¸   r¹   rÎ   rÏ   rÐ   rÑ   c              3   ó@   #   • U  H  u  pp4[        XU5      v •  M     g 7fr   r    rÙ   s        r!   r"   Ú4_efficient_attention_forward_flop.<locals>.<genexpr>æ  rÛ   rÜ   ©rÖ   rÞ   )r·   r¸   r¹   ÚbiasrÎ   rÏ   rÐ   rÑ   r2   r3   rß   s              r!   Ú!_efficient_attention_forward_flopræ   Ì  s?   € ô" 6ØØØØ!Ø!Ø!Ø!ñ€Eô ñ á6;óó ð r,   c                 óØ  • SnUu  pVpxUu  pšp¼Uu  pÞnnU u  nnnnXYs=:X  a  Us=:X  a  U:X  a  O  OXjs=:X  a  Us=:X  a  U:X  a  O  OXŒ:X  d  [        S5      eUU:X  a  X¿:X  a  UU:X  d  [        S5      eSnU[        XV-  Xx4XV-  X‹45      -  nU[        XV-  UU4XV-  UU45      -  nU[        XV-  X·4XV-  UU45      -  nU[        XV-  X{4XV-  X¸45      -  nU[        XV-  X‡4XV-  X{45      -  nU$ )Nr   zFsdpa_backward_flop_count: batch/heads/dimension mismatch among tensorszJsdpa_backward_flop_count: grad_out/value/key/query shapes are incompatibler�   )r„   rŽ   r�   r�   rœ   r`   r‘   r’   r“   r”   r•   r–   r—   r˜   r™   rš   r›   Ú_b4Ú_h4Ú_s4Ú_d4s                        r!   Úsdpa_backward_flop_countrì   ì  s2  € Ø€KØ �N€Aˆ#Ø"Ñ€CˆcØ$Ñ€Cˆc�3Ø'Ñ€Cˆˆc�3ØÕ!�sÕ!˜cÖ!¨Õ)?°SÕ)?¸CÖ)?ÀsÃzÜÐeÓfÐfØ�#‹:˜S›Z¨s°c«zÜÐiÓjÐjØ€Kð ”8˜Q™U CÐ-°±°sÐ/@ÓAÑA€Kð ”8˜Q™U C¨Ð-°±°s¸CÐ/@ÓAÑA€Kà”8˜Q™U CÐ-°±°s¸CÐ/@ÓAÑA€Kð ”8˜Q™U CÐ-°±°sÐ/@ÓAÑA€Kà”8˜Q™U CÐ-°±°sÐ/@ÓAÑA€KØÐr,   c                ó   • [        XX#5      $ )z(Count flops for self-attention backward.©rì   )r„   rŽ   r�   r�   r1   r2   r3   s          r!   Úsdpa_backward_floprï   	  s   € ô
 $ NÀÓXÐXr,   c
                 óF   • [        UUUU UUUU	S9n[        S U 5       5      $ )N)r·   r¸   r¹   r¯   rº   r»   r¼   r½   c              3   ó@   #   • U  H  u  pp4[        XAX#5      v •  M     g 7fr   rî   ©r   rŽ   r�   r�   r„   s        r!   r"   Ú1_flash_attention_backward_flop.<locals>.<genexpr>+  ó&   é € ð âCIÑ?ˆK Kô 	! ¸i×UÐUÚCIùrÜ   rÝ   )r¯   r·   r¸   r¹   ÚoutÚ	logsumexprº   r»   r¼   r½   r2   r3   Úshapess                r!   Ú_flash_attention_backward_floprø     sB   € ô" 3ØØØØØØØØñ	€Fô ñ áCIóó ð r,   c
                 óF   • [        UUUU UUUU	S9n[        S U 5       5      $ )N)r·   r¸   r¹   r¯   rÎ   rÏ   rÐ   rÑ   c              3   ó@   #   • U  H  u  pp4[        XAX#5      v •  M     g 7fr   rî   rò   s        r!   r"   Ú5_efficient_attention_backward_flop.<locals>.<genexpr>L  rô   rÜ   rä   )r¯   r·   r¸   r¹   rå   rõ   rÎ   rÏ   rÐ   rÑ   r2   r3   r÷   s                r!   Ú"_efficient_attention_backward_floprü   1  sB   € ô" 7ØØØØØ!Ø!Ø!Ø!ñ	€Fô ñ áCIóó ð r,   c                 ó6   • [        U [        5      (       d  U 4$ U $ r   )r'   Útuple)Úxs    r!   Únormalize_tupler   j  s   € Ü�aœ×ÑØˆtˆØ€Hr,   )Ú ÚKÚMÚBÚTc                 ó�   • [        S[        [        [        5      S-
  [        [	        U 5      5      S-
  S-  5      5      n[        U   $ )Nr   r   rP   r±   )ÚmaxÚminrµ   ÚsuffixesÚstr)ÚnumberÚindexs     r!   Úget_suffix_strr  s  s=   € ô �”3”sœ8“} qÑ(¬3¬s°6«{Ó+;¸aÑ+?ÀAÑ*EÓFÓG€EÜ�E‰?Ðr,   c                 óX   • [         R                  U5      nU SU-  -  S nU[         U   -   $ )Niè  z.3f)r	  r  )r  Úsuffixr  r¹   s       r!   Úconvert_num_with_suffixr  z  s2   € Ü�N‰N˜6Ó"€Eà˜ ™Ñ% cÐ*€Eà”8˜E‘?Ñ"Ð"r,   c                 ó   • US:X  a  gX-  S $ )Nr   ú0%z.2%© )ÚnumÚdenoms     r!   Úconvert_to_percent_strr  �  s   € Ø�ƒzØØ‰k˜#ÐÐr,   c                 ó0   ^ • [        T 5      U 4S j5       nU$ )Nc                 ó>   >• [        U 5      u  pT" U6 n[        X25      $ r   )r   r   )r2   Ú	flat_argsÚspecrõ   r4   s       €r!   r5   Ú)_pytreeify_preserve_structure.<locals>.nf‡  s#   ø€ ä& tÓ,‰ˆ	Ù�ˆmˆÜ˜cÓ(Ð(r,   r   r7   s   ` r!   Ú_pytreeify_preserve_structurer  †  s    ø€ Ü
ˆ1ƒXô)ó ð)ð
 €Ir,   c                   ó  ^ • \ rS rSrSr    SS\R                  R                  \\R                  R                     -  S-  S\	S\
S\\\4   S-  SS4
U 4S	 jjjrS\	4S
 jrS\\\\\	4   4   4S jrSS jrS rS rS rSrU =r$ )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()

NÚmodsÚdepthÚdisplayÚcustom_mappingr9   c                 ón  >• [         TU ]  5         [        S 5      U l        X l        X0l        S U l        Uc  0 nUb  [        R                  " SSS9  0 [        EUR                  5        VVs0 s H%  u  pVU[        USS5      (       a  UO
[        U5      _M'     snnEU l	        [        5       U l        g s  snnf )Nc                  ó    • [        [        5      $ r   )r   Úintr  r,   r!   Ú<lambda>Ú*FlopCounterMode.__init__.<locals>.<lambda>«  s
   € Ì+ÔVYÔJZr,   z<mods argument is not needed anymore, you can stop passing itrP   )Ú
stacklevelÚ_get_rawF)ÚsuperÚ__init__r   Úflop_countsr  r   ÚmodeÚwarningsÚwarnr-   Úitemsr   r8   r   Úmod_tracker)Úselfr  r  r   r!  rV   ÚvÚ	__class__s          €r!   r*  ÚFlopCounterMode.__init__¤  s²   ø€ ô 	‰ÑÔÜ6AÑBZÓ6[ˆÔØŒ
ØŒØ-1ˆŒ	ØÑ!ØˆNØÑÜ�MŠMÐXÐefÒgð
Üð
àWe×WkÑWkÔWmÔnÒWmÉtÈqˆq”w˜q *¨e×4Ñ4‘!¼-ÈÓ:JÒJÑWmÒnð
ˆÔô )›?ˆÕùó os   Á+,B1c                 óN   • [        U R                  S   R                  5       5      $ )NÚGlobal)rÞ   r+  Úvalues©r1  s    r!   Úget_total_flopsÚFlopCounterMode.get_total_flops¹  s!   € Ü�4×#Ñ# HÑ-×4Ñ4Ó6Ó7Ð7r,   c                 ó€   • U R                   R                  5        VVs0 s H  u  pU[        U5      _M     snn$ s  snnf )zæReturn 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/  Údict)r1  rV   r2  s      r!   Úget_flop_countsÚFlopCounterMode.get_flop_counts¼  s7   € ð (,×'7Ñ'7×'=Ñ'=Ô'?Ô@Ò'?™t˜q�”4˜“7’
Ñ'?Ò@Ð@ùÓ@s   ž:c                 ó(  ^ ^
^^• Uc  T R                   nUc  SnSS KnSUl        / SQn/ nT R                  5       m
[	        T
5      mSmU
UUU 4S jn[        T R                  R                  5       5       HB  nUS:X  a  M  UR                  S5      S	-   nXq:”  a  M&  U" XgS	-
  5      nUR                  U5        MD     ST R                  ;   a'  T(       d   U H  n	S
U	S   -   U	S'   M     U" SS5      U-   n[        U5      S:X  a  / SQ/nUR                  XCSS9$ )Ni?B r   T)ÚModuleÚFLOPz% TotalFc           	      ó€  >• [        T
R                  U    R                  5       5      nT	UT:¬  -  m	SU-  n/ nUR                  X0-   [	        UT5      [        UT5      /5        T
R                  U    R                  5        H<  u  pVUR                  US-   [        U5      -   [	        UT5      [        UT5      /5        M>     U$ )NÚ z - )rÞ   r+  r7  Úappendr  r  r/  r
  )Úmod_namer  rœ   Úpaddingr7  rV   r2  Úglobal_flopsÚglobal_suffixÚis_global_subsumedr1  s          €€€€r!   Úprocess_modÚ.FlopCounterMode.get_table.<locals>.process_modØ  sÊ   ø€ ô ˜d×.Ñ.¨xÑ8×?Ñ?ÓAÓBˆKà +°Ñ"=Ñ=Ðà˜E‘kˆGØˆFØ�M‰MØÑ"Ü'¨°]ÓCÜ& {°LÓAðô ð
 ×(Ñ(¨Ñ2×8Ñ8Ö:‘�Ø—‘Ø˜e‘O¤c¨!£fÑ,Ü+¨A¨}Ó=Ü*¨1¨lÓ;ðö ñ ;ð ˆMr,   r6  Ú.r   rC  )r6  Ú0r  )ÚleftÚrightrO  )ÚheadersÚcolalign)r  ÚtabulateÚPRESERVE_WHITESPACEr9  r  Úsortedr+  ÚkeysÚcountÚextendrµ   )r1  r  rR  Úheaderr7  rJ  ÚmodÚ	mod_depthÚ
cur_valuesr¹   rG  rH  rI  s   `         @@@r!   Ú	get_tableÚFlopCounterMode.get_tableÈ  s%  û€ Ø‰=Ø—J‘JˆEØ‰=ØˆEó 	à'+ˆÔ$Ú.ˆØˆØ×+Ñ+Ó-ˆÜ& |Ó4ˆØ"Ð÷	ð 	ô, ˜$×*Ñ*×/Ñ/Ó1Ö2ˆCØ�h‹ÙØŸ	™	 #›¨Ñ*ˆIØÓ Ùá$ S°a©-Ó8ˆJØ�M‰M˜*Ö%ñ 3ð �t×'Ñ'Ó'Ö0BÛ�Ø  q¡™>��a“ñ  ñ ! ¨1Ó-°Ñ6ˆFäˆv‹;˜!ÓÚ+Ð,ˆFà× Ñ  ÐB\Ð Ð]Ð]r,   c                 óÂ   • U R                   R                  5         U R                  R                  5         [	        U 5      U l        U R
                  R                  5         U $ r   )r+  Úclearr0  Ú	__enter__Ú_FlopCounterModer,  r8  s    r!   r`  ÚFlopCounterMode.__enter__  sG   € Ø×Ñ×ÑÔ Ø×Ñ×"Ñ"Ô$Ü$ TÓ*ˆŒ	Ø�	‰	×ÑÔØˆr,   c                 ó  • U R                   c  [        S5      eU R                   R                  " U6 nS U l         U R                  R                  5         U R                  (       a$  [        U R                  U R                  5      5        U$ )Nz<Internal error: FlopCounter.__exit__ called but mode is None)r,  rR   Ú__exit__r0  r   Úprintr\  r  )r1  r2   r`   s      r!   rd  ÚFlopCounterMode.__exit__  sf   € Ø�9‰9ÑÜ Ð!_Ó`Ð`Ø�I‰I×Ò Ð%ˆØˆŒ	Ø×Ñ×!Ñ!Ô#Ø�<�<Ü�$—.‘. §¡Ó,Ô-Øˆr,   c                 óÚ   • XR                   ;   a[  U R                   U   nU" U0 UDSU0D6n[        U R                  R                  5       H  nU R                  U   U==   U-  ss'   M     U$ )Nr/   )r-   Úsetr0  Úparentsr+  )r1  Úfunc_packetrõ   r2   r3   Úflop_count_funcrˆ   Úpars           r!   Ú_count_flopsÚFlopCounterMode._count_flops  sm   € Ø×,Ñ,Ó,Ø"×0Ñ0°Ñ=ˆOÙ(¨$ÐF°&ÑFÀ#ÒFˆJÜ˜4×+Ñ+×3Ñ3Ö4�Ø× Ñ  Ñ% kÓ2°jÑ@Õ2ñ 5àˆ
r,   )r  r   r+  r-   r0  r,  )NrP   TNr   )Ú__name__Ú
__module__Ú__qualname__Ú__firstlineno__Ú__doc__r   Únnr@  r�   r$  Úboolr<  r	   r*  r9  r
  r=  r\  r`  rd  rm  Ú__static_attributes__Ú__classcell__)r3  s   @r!   r   r   �  sÆ   ø† ñð* DHØØ Ø48ñ+à—(‘(—/‘/ D¨¯©¯©Ñ$9Ñ9¸DÑ@ð+ð ð+ð ð	+ð
 !  c ™N¨TÑ1ð+ð
 >B÷+ð +ð*8 ô 8ð
A  c¨4°°S°©>Ð&9Ñ!:ô 
Aô<^ò~ò÷ð r,   c                   ó@   • \ rS rSrSrS\SS4S jrS rS rSS	 jr	S
r
g)ra  i   TÚcounterr9   Nc                 ó   • Xl         g r   ©ry  )r1  ry  s     r!   r*  Ú_FlopCounterMode.__init__#  s   € Ø�r,   c                 ó  • SSK nUR                  U R                  R                  5      nU    U" U6 nSSS5        UR                  U R                  R                  5      nX@R                  l        WU4$ ! , (       d  f       NG= f)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
r   N)Úcopyry  r+  )r1  Ú	branch_fnÚoperandsr~  Úcheckpointed_flop_countsÚresultr+  s          r!   Ú$_execute_with_isolated_flop_countingÚ5_FlopCounterMode._execute_with_isolated_flop_counting&  sg   € ó 	Ø#'§9¡9¨T¯\©\×-EÑ-EÓ#FÐ ÚÙ Ð)ˆF÷ à—i‘i §¡× 8Ñ 8Ó9ˆØ#;�‰Ô Ø�{Ð"Ð"÷	 �Tús   ¬A3Á3
Bc                 ó„  • U[         R                  R                  R                  [         R                  R                  R                  1;   nU(       au  SSKJn  SSKJn  U" US   5      n[        X‡5      (       d1  [        US5      (       a  UR                  nOO[        X‡5      (       d  M1  U R                  R                  US X45      $ U[         R                  R                  R                  L GaL  Uu  pšp¼U R                  X¬5      u  pÞU[         L a  [         $ U R                  X¼5      u  nnU[         L a  [         $ [#        UR%                  5       5      [#        UR%                  5       5      -  n0 nU Hƒ  nUU   nUU   n0 n[#        UR%                  5       5      [#        UR%                  5       5      -  nU H6  nUR'                  US5      nUR'                  US5      n[)        UU5      UU'   M8     UUU'   M…     UR+                  5        H.  u  nnU R                  R,                  U   R/                  U5        M0     U$ [         $ )Nr   )Ú
get_kernelr   Ú
kernel_idxÚfn)r   ÚopsÚhigher_orderÚtriton_kernel_wrapper_mutationÚ triton_kernel_wrapper_functionalÚ*torch._higher_order_ops.triton_kernel_wrapr†  Útriton.runtime.jitr   r'   Úhasattrrˆ  ry  rm  Úcondrƒ  ÚNotImplementedrh  rU  Úgetr  r/  r+  Úupdate)r1  ÚfuncÚtypesr2   r3   Ú	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   € ØœUŸY™Y×3Ñ3×RÑRÜ"ŸY™Y×3Ñ3×TÑTðVñ Vˆ	æÝMå6Ù$ V¨LÑ%9Ó:ˆKä  ×:Ñ:Ü˜;¨×-Ñ-Ø"-§.¡.‘Kàô	 ! ×:Ó:ð
 —<‘<×,Ñ,¨[¸$ÀÓMÐMØ”U—Y‘Y×+Ñ+×0Ñ0Ó0ð
 9=Ñ5ˆD˜|à)-×)RÑ)RØó*Ñ&ˆHð œ>Ò)Ü%Ð%à+/×+TÑ+TØó,Ñ(ˆIÐ(ð œNÒ*Ü%Ð%ô Ð/×4Ñ4Ó6Ó7¼#Ð>O×>TÑ>TÓ>VÓ:WÑWˆLØ!#ÐÛ)�	Ø#3°IÑ#>Ð Ø$5°iÑ$@Ð!à%'Ð"Ü #Ð$4×$9Ñ$9Ó$;Ó <¼sÐCT×CYÑCYÓC[Ó?\Ñ \�ã -�HØ/×3Ñ3°H¸aÓ@�HØ 1× 5Ñ 5°hÀÓ B�IÜ36°xÀÓ3KÐ& xÓ0ñ !.ð
 1CÐ" 9Ó-ñ *ð *<×)AÑ)AÖ)CÑ%�	˜:Ø—‘×(Ñ(¨Ñ3×:Ñ:¸:ÖFñ *Dð
 ˆOä!Ð!r,   c                 ód  • U(       a  UO0 nU[         R                  R                  R                  R                  [         R                  R                  R
                  R                  [         R                  R                  R
                  R                  [         R                  R                  R                  R                  [         R                  R                  R                  R                  [         R                  R                  R                  R                  [         R                  R                  R                  R                  [         R                  R                  R                  R                  [         R                  R                  R                  R                  [         R                  R                  R                  R                  [         R                  R                  R                  R                  [         R                  R                  R                  R                  [         R                  R                  R                   R                  [         R                  R                  R"                  R                  [         R                  R$                  R&                  R                  1;   a  [(        $ [+        U[         R,                  R.                  5      (       a  U R1                  XX45      $ XR2                  R4                  ;  ac  U[         R                  R$                  R6                  R                  La2  U    UR8                  " U0 UD6nU[(        La  UsS S S 5        $  S S S 5        U" U0 UD6nU R2                  R;                  UR<                  XcU5      $ ! , (       d  f       N== f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'   r>   ÚHigherOrderOperatorrª  ry  r-   r¨   Ú	decomposerm  Ú_overloadpacket)r1  r”  r•  r2   r3   Úrrõ   s          r!   Ú__torch_dispatch__Ú#_FlopCounterMode.__torch_dispatch__w  s9  € Þ!‘ rˆð ”E—I‘I—N‘N×4Ñ4×<Ñ<Ü—I‘I—N‘N×0Ñ0×8Ñ8Ü—I‘I—N‘N×0Ñ0×>Ñ>Ü—I‘I—N‘N×9Ñ9×AÑAÜ—I‘I—N‘N×?Ñ?×GÑGÜ—I‘I—N‘N×'Ñ'×/Ñ/Ü—I‘I—N‘N×+Ñ+×3Ñ3Ü—I‘I—N‘N×)Ñ)×1Ñ1Ü—I‘I—N‘N×-Ñ-×5Ñ5Ü—I‘I—N‘N×1Ñ1×9Ñ9Ü—I‘I—N‘N×5Ñ5×=Ñ=Ü—I‘I—N‘N×(Ñ(×0Ñ0Ü—I‘I—N‘N×,Ñ,×4Ñ4Ü—I‘I—N‘N×&Ñ&×.Ñ.Ü—I‘I—N‘N×)Ñ)×1Ñ1ð3ó 3ô  "Ð!ä�dœEŸJ™J×:Ñ:×;Ñ;Ø×0Ñ0°¸dÓKÐKð —|‘|×1Ñ1Ó1°dÄ%Ç)Á)Ç.Á.×BWÑBW×B_ÑB_Ò6_ÚØ—N’N DÐ3¨FÑ3�ØœNÒ*Ø÷ ‘à*÷ ñ �DÐ#˜FÑ#ˆØ�|‰|×(Ñ(¨×)=Ñ)=¸sÈ&ÓQÐQ÷ •ús   ÍN!Î!
N/r{  )r  N)ro  rp  rq  rr  Úsupports_higher_order_operatorsr   r*  rƒ  rª  rÂ  rv  r  r,   r!   ra  ra     s,   † Ø&*Ð#ð ð °Dô ò#ò(;"÷z"Rr,   ra  )Fr   )NNNFN)dr•  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   Ú	getLoggerro  ÚlogrŽ  r   r@   ÚImportErrorÚanyÚwarningr‰  r­  r+   r-   r<  Ú__annotations__r8   r   Úmmr$  rY   Úaddmmr^   Úbmmrc   Úbaddbmmre   Ú
_scaled_mmrm   r�   ru  rw   Ú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Ë   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  r  r
  r  r  r   ra  r  r,   r!   Ú<module>rë     s?  ðæ Û Û ß FÑ FÝ )ß Ý $Ý $Ý 'Ý #Ý :Ý Ý Û àÐ5Ð
6€áˆTƒ]€Ùˆtƒ_€à×Ò˜Ó!€ðÝ>ð ‡y�y‡~�~€òð
 !#€ˆt�C˜�H‰~Ó "òñ°X¸xÈÈBÈÑ?OÐ>PÐRZÐ[]Ð_aÐ[aÑRbÐ>bÑ5cõ ñ. �t—w‘wÓØ/3ò 	À#ô 	ó  ð	ñ �t—z‘zÓ"ñ%È#ô %ó #ð%ñ �t—x‘xÓ ñ¸Cô ó !ðñ �t—|‘|Ó$ñ&ÈCô &ó %ð&ñ �t—‘Ó'ð ØØØØñ%ð 	ô%ó (ð%ð( ñ	$Ø�#‰Yð$à�#‰Yð$ð �C‰yð$ð ð	$ð
 	õ$ñL ˜×(Ñ(Ø×)Ñ)Ø×.Ñ.Ø×1Ñ1Ø×5Ñ5ð	7ó 8ð
 cgò OÐuxô Oó8ð
Oñ �t×0Ñ0Ó1ðeð óeó 2ðeòNñ& ˜×DÑDØ×@Ñ@Ø×@Ñ@ðBó Cð EIò @ÐWZô @óCð@ò	-ð" ò1`ð ˆe�E˜#˜s˜(‘O U¨3°¨8¡_°e¸CÀ¸H±oÀuÈSÐRUÈXÁÐY]ÑG]Ð]Ñ^Ñ_õ1`ðr ò4`ð ˆe�E˜#˜s˜(‘O U¨3°¨8¡_°e¸CÀ¸H±oÀuÈSÐRUÈXÁÐY]ÑG]Ð]Ñ^Ñ_õ4`ñn �t×4Ñ4¸dÑCð òð 	ôó Dðñ> �t×8Ñ8À$ÑGðð 	óó Hðò>ñ: ˜×MÑMØ×IÑIØ×IÑIðKó Lð ^bò YÐpsô YóLðYñ �t×5Ñ5¸tÑDðð 	óó Eðñ@ �t×9Ñ9À4ÑHðð 	óó Iðð@Ø‡G�GˆWðà‡J�J�
ðð 	‡H�Hˆhðð 	‡L�L�,ð	ð
 	‡O�O�_ðð 	×Ñ�iðð 	×Ñ�yðð 	×Ñ˜Iðð 	×!Ñ! 9ðð 	×Ñ˜yðð 	×ÑÐ1ðð 	×0Ñ0°)ðð 	×,Ñ,¨iðð 	×,Ñ,¨iðð 	×9Ñ9Ð;Mðð  	×5Ñ5Ð7Ið!ð" 	×5Ñ5Ð7Ið#ð$ 	×!Ñ!Ð#@Ø×%Ñ%Ð'HØ×"Ñ"Ð$BØ×&Ñ&Ð(Jñ+€ò0ò $€òò#ð ¨#ô  ò
÷Nñ Nô`yRÐ(õ yRøðK ó Ù
Ñ
]ÑF\Ó
]×]Ñ]Ø�‰ÐVÔWØƒLðús   Á=Q" Ñ"-RÒR