Ë
    ùÿæiòŠ  ã                   óú  — U d dl mZ d dlZd dlZd dlmZmZmZ ddl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 d dlmZ d dlmZ d dlmZ d dlZddgZ ed«      Z ed«      Z ej<                  e«      Z 	 d dl!m"Z# ejN                  jP                  Z(d„ Z)i Z*e+eef   e,d<   d„ Z-d@deeeef   geeef   f   fd„Z. e.e(j^                  «      ddœde0fd„«       Z1 e.e(jd                  «      dAde0fd„«       Z3 e.e(jh                  «      dAde0fd„«       Z5 e.e(jl                  «      dAde0fd„«       Z7 e.e(jp                  «      	 	 	 	 	 dBde0fd „«       Z9	 d@d!e:e0   d"e:e0   d#e:e0   d$e;de0f
d%„Z< e.e(jz                  e(j|                  e(j~                  e(j€                  e(j‚                  g«      ddœde0fd&„«       ZB e.e(j†                  «      de0fd'„«       ZDd(„ ZE e.e(jŒ                  e(jŽ                  e(j�                  g«      ddœde0fd)„«       ZId*„ ZJdd+œdeeKeKe0d,f   eKe0d,f   eKe0d,f   eKe0d,f   dz  f      fd-„ZLdd+œdeeKeKe0d,f   eKe0d,f   eKe0d,f   eKe0d,f   dz  f      fd.„ZM e.e(jœ                  d/¬0«      ddœde0fd1„«       ZO e.e(j                   d/¬0«      de0fd2„«       ZQd3„ ZR e.e(j¦                  e(j¨                  e(jª                  g«      ddœde0fd4„«       ZV e.e(j®                  d/¬0«      de0fd5„«       ZX e.e(j²                  d/¬0«      de0fd6„«       ZZi e(j^                  e1“e(jd                  e3“e(jh                  e5“e(jl                  e7“e(jp                  e9“e(jz                  eB“e(j|                  eB“e(j~                  eB“e(j‚                  eB“e(j€                  eB“e(j†                  eD“e(jŒ                  eI“e(jŽ                  eI“e(j�                  eI“e(j¦                  eV“e(j¨                  eV“e(jª                  eV“e(jœ                  eOe(j                   eQe(j®                  eXe(j²                  eZi¥Z*d7„ Z[g d8¢Z\d9„ Z]d:„ Z^de_fd;„Z`d<„ Za G d=„ d«      Zb G d>„ d?e«      Zcy# e$$ r&  e%d„ dD «       «      re jM                  d«       eZ#Y �Œöw xY w)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              #   óV   K  — | ]!  }t        t        j                  |d «      d u–— Œ# y ­w©N)ÚgetattrÚtorchÚversion)Ú.0Úattrs     úm/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/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                 óR   — t        | t        j                  «      r| j                  S | S r   )Ú
isinstancer   ÚTensorÚshape)Úis    r!   Ú	get_shaper+   #   s   € Ü�!”U—\‘\Ô"Ø�w‰wˆØ€Hó    Úflop_registryc                 ó4   ‡ — t        ‰ «      d dœˆ fd„
«       }|S )N)Úout_valc                 óF   •— t        t        ||| f«      \  }}} ‰|d|i|¤ŽS )NÚ	out_shape)r   r+   )r/   ÚargsÚkwargsr1   Úfs       €r!   Únfzshape_wrapper.<locals>.nf+   s2   ø€ ä"*¬9°t¸VÀWÐ6MÓ"NÑˆˆf�iÙ�$Ð6 )Ð6¨vÑ6Ð6r,   r   ©r4   r5   s   ` r!   Úshape_wrapperr7   *   s#   ø€ Ü
ˆ1ƒXØõ 7ó ð7ð €Ir,   Úreturnc                 ód   ‡ ‡— dt         t        t        f   dt         t        t        f   fˆˆ fd„}|S )NÚflop_formular8   c                 ó„   •‡ — ‰st        ‰ «      Š dˆ fd„}t        j                  j                  j	                  |‰«       ‰ S )Nc                 óÌ   •— t        | t        j                  j                  t        f«      st        d| › dt        | «      › �«      ‚| t        v rt        d| › �«      ‚‰t        | <   y )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!   Úregisterz=register_flop_formula.<locals>.register_fun.<locals>.register7   sm   ø€ Ü˜v¬¯
©
×(CÑ(CÄ\Ð'RÔSÜ ðà#˜HÐ$6´t¸F³|°nðFóGð Gð œÑ&Ü"Ð%AÀ&ÀÐ#JÓKÐKØ$0ŒM˜&Ò!r,   )r8   N)r7   r   ÚutilsÚ_pytreeÚ	tree_map_)r:   rD   Úget_rawÚtargetss   ` €€r!   Úregister_funz+register_flop_formula.<locals>.register_fun3   s7   ù€ ÙÜ(¨Ó6ˆLõ	1ô 	�‰×Ñ×%Ñ% h°Ô8àÐr,   )r   r   r   )rI   rH   rJ   s   `` r!   r   r   1   s0   ù€ ð¤8¬B´¨FÑ#3ð ¼ÄÄRÀÑ8Hö ð& Ðr,   )r1   c                óX   — | \  }}|\  }}||k7  rt        d|› d|› �«      ‚||z  dz  |z  S )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_floprV   H   sM   € ð
 �D€A€qØ�E€BˆØˆB‚wÜÐRÐSTÐRUÐUZÐ[]ÐZ^Ð_Ó`Ð`àˆq‰5�1‰9�q‰=Ðr,   c                 ó   — t        ||«      S )zCount flops for addmm.©rV   ©Ú
self_shaperP   rQ   r1   r3   s        r!   Ú
addmm_flopr[   T   s   € ô �7˜GÓ$Ð$r,   c                 ó’   — | \  }}}|\  }}}	||k7  rt        d|› d|› �«      ‚||k7  rt        d|› d|› �«      ‚||z  |	z  dz  |z  }
|
S )z"Count flops for the bmm operation.z0bmm: batch dimensions must match (b == b2), got rL   z0bmm: inner dimensions must match (k == k2), got rM   rN   )rP   rQ   r1   r3   ÚbrR   rS   Úb2rT   rU   Úflops              r!   Úbmm_flopr`   Y   s}   € ð
 �G€A€qˆ!Ø�I€BˆˆAØˆB‚wÜÐOÐPQÈsÐRWÐXZÐW[Ð\Ó]Ð]ØˆB‚wÜÐOÐPQÈsÐRWÐXZÐW[Ð\Ó]Ð]àˆq‰5�1‰9�q‰=˜1Ñ€DØ€Kr,   c                 ó   — t        ||«      S )z&Count flops for the baddbmm operation.)r`   rY   s        r!   Úbaddbmm_floprb   h   s   € ô
 �G˜WÓ%Ð%r,   c	                 ó   — t        | |«      S )zCount flops for _scaled_mm.rX   )
rP   rQ   Úscale_a_shapeÚscale_b_shapeÚ
bias_shapeÚscale_result_shapeÚ	out_dtypeÚuse_fast_accumr1   r3   s
             r!   Ú_scaled_mm_floprj   o   s   € ô �7˜GÓ$Ð$r,   Úx_shapeÚw_shaper1   Ú
transposedc                 ót   — | d   }|r| n|dd }|^}}}	 t        |«      t        |«      z  |z  |z  |z  dz  }	|	S )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   rM   Nr   )
rk   rl   r1   rm   Ú
batch_sizeÚ
conv_shapeÚc_outÚc_inÚfilter_sizer_   s
             r!   Úconv_flop_countrt   €   s]   € ð( ˜‘€JÙ'‘'¨Y¸¸Ð;€JØ 'Ð€Eˆ4�+ðô �
Óœd ;Ó/Ñ/°*Ñ<¸uÑDÀtÑKÈaÑO€DØ€Kr,   c                ó    — t        | |||¬«      S )zCount flops for convolution.©rm   )rt   )
rk   rl   Ú_biasÚ_strideÚ_paddingÚ	_dilationrm   r1   r2   r3   s
             r!   Ú	conv_flopr{   ¦   s   € ô ˜7 G¨YÀ:ÔNÐNr,   c                 ó  — d„ }d}	 |
d   r t        |d   «      }|t        | ||| «      z  }|
d   rZt        |d   «      }|r&|t         || «       ||«       ||«      d¬«      z  }|S |t         ||«       || «       ||«      d¬«      z  }|S )Nc                 ó4   — | d   | d   gt        | dd  «      z   S )Nr   r   rM   )Úlist)r)   s    r!   Útzconv_backward_flop.<locals>.tÀ   s$   € Ø�a‘˜% ™(Ð#¤d¨5°°¨9£oÑ5Ð5r,   r   r   Frv   )r+   rt   )Úgrad_out_shaperk   rl   rw   rx   ry   rz   rm   Ú_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°gÐ?OÐU_Ð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                 ó4  — | \  }}}}|\  }}}	}
|\  }}}}||cxk(  r|k(  r4n t        d«      ‚||cxk(  r|k(  rn t        d«      ‚||
k(  r
|	|k(  r||
k(  st        d«      ‚d}|t        ||z  ||f||z  ||	f«      z  }|t        ||z  ||	f||z  |	|f«      z  }|S )z^
    Count flops for self-attention.

    NB: We can assume that value_shape == key_shape
    z8sdpa_flop_count: query/key/value shapes are incompatibler   ©rO   r`   )Ú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€qˆ#ˆsØ"Ñ€Cˆˆc�3Ø$Ñ€Cˆˆc�3Ø�Œ?�sŒ?ÜÐWÓXÐXð #$ s¤/¨c¤/ÜÐWÓXÐXð :=ÀºÈ3ÐRUÊ:Ð]`ÐdgÒ]gÜÐWÓXÐXØ€Kà”8˜Q ™U C¨Ð-°°A±°s¸CÐ/@ÓAÑA€Kà”8˜Q ™U C¨Ð-°°A±°s¸CÐ/@ÓAÑA€KØÐr,   c                ó   — t        | ||«      S )úCount flops for self-attention.©r™   )rŠ   r‹   rŒ   r1   r2   r3   s         r!   Ú	sdpa_flopr�   ,  s   € ô ˜;¨	°;Ó?Ð?r,   c                 óÔ   — ddl m} ddlm} t	        | ||f«      s7| j
                  j                  dk7  r| j                  «       j                  «       S |g| j                  d«      dz
  z  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'   ÚdevicerA   Ú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              #   óÌ  K  — |��)t        |j                  «      dk7  rt        d«      ‚t        |j                  «      dk7  rt        d«      ‚|�$|j                  | j                  k7  rt        d«      ‚| j                  \  }}	}
|j                  \  }}}|j                  \  }}}|€t        d«      ‚|€t        d«      ‚|j                  |j                  k7  rt        d«      ‚t        ||«      }t        ||«      }t	        ||d	¬
«      D ]%  \  }}d|	||
f}d|||f}d|||f}|�|nd}||||f–— Œ' y| j                  |j                  |j                  |�|j                  ndf–— y­w)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)   rO   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  s{  è ø€ ð$ Ñô ˆs�y‰y‹>˜QÒÜ Ð!ZÓ[Ð[Üˆu�{‰{Ó˜qÒ Ü Ð!\Ó]Ð]ØÐ H§N¡N°e·k±kÒ$AÜ Ð!gÓhÐhØ—k‘k‰ˆˆ3�Ø—i‘i‰ˆˆ3�Ø—k‘k‰ˆˆ3�ØÐÜ Ð!NÓOÐOØÐÜ Ð!NÓOÐOØ�?‰?˜iŸo™oÒ-Ü Ð!dÓeÐeÜ+¨I°uÓ=ˆÜ+¨I°uÓ=ˆÜ&)¨-¸È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              #   óÒ  K  — |��,t        |j                  «      dk7  rt        d«      ‚t        |j                  «      dk7  rt        d«      ‚|�$|j                  | j                  k7  rt        d«      ‚| j                  \  }}}	}
|j                  \  }}}}|j                  \  }}}}|€t        d«      ‚|€t        d«      ‚|j                  |j                  k7  rt        d«      ‚t        ||«      }t        ||«      }t	        ||d	¬
«      D ]%  \  }}d|	||
f}d|||f}d|||f}|�|nd}||||f–— Œ' y| j                  |j                  |j                  |�|j                  ndf–— y­w)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ˆc�3ØŸ™‰ˆˆ1ˆc�3ØŸ™‰ˆˆ1ˆc�3ØÐÜ Ð!kÓlÐlØÐÜ Ð!kÓlÐlØ×Ñ ×!3Ñ!3Ò3Ü ð "Zó [ð [ä'¨°lÓCˆ	Ü'¨°lÓ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)rH   c          	      óJ   — t        | ||||||¬«      }
t        d„ |
D «       «      S )r›   )r³   r´   rµ   r¶   r·   r¸   r¹   c              3   ó@   K  — | ]  \  }}}}t        |||«      –— Œ y ­wr   rœ   ©r   rŠ   r‹   rŒ   rº   s        r!   r"   z0_flash_attention_forward_flop.<locals>.<genexpr>Æ  ó*   è ø€ ð á6;Ñ2ˆK˜ K°ô 	˜ Y°×<Ù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           	      óJ   — t        | ||||||¬«      }
t        d„ |
D «       «      S )r›   )r³   r´   rµ   rÊ   rË   rÌ   rÍ   c              3   ó@   K  — | ]  \  }}}}t        |||«      –— Œ y ­wr   rœ   rÕ   s        r!   r"   z4_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                 ó   — d}|\  }}}}|\  }	}
}}|\  }}}}| \  }}}}||	cxk(  r|cxk(  r|k(  r0n t        d«      ‚||
cxk(  r|cxk(  r|k(  rn t        d«      ‚||k(  st        d«      ‚||k(  r
||k(  r||k(  st        d«      ‚d}|t        ||z  ||f||z  ||f«      z  }|t        ||z  ||f||z  ||f«      z  }|t        ||z  ||f||z  ||f«      z  }|t        ||z  ||f||z  ||f«      z  }|t        ||z  ||f||z  ||f«      z  }|S )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æ   ì  s†  € Ø€KØ �N€A€qˆ#ˆsØ"Ñ€Cˆˆc�3Ø$Ñ€Cˆˆc�3Ø'Ñ€Cˆˆc�3Ø�Ô!�sÔ!˜cÔ!ÜÐeÓfÐfð *+¨cÔ)?°SÔ)?¸CÔ)?ÜÐeÓfÐfð HKÈcÂzÜÐeÓfÐfØ�#Š:˜S CšZ¨s°cªzÜÐiÓjÐjØ€Kð ”8˜Q ™U C¨Ð-°°A±°s¸CÐ/@ÓAÑA€Kð ”8˜Q ™U C¨Ð-°°A±°s¸CÐ/@ÓAÑA€Kà”8˜Q ™U C¨Ð-°°A±°s¸CÐ/@ÓAÑA€Kð ”8˜Q ™U C¨Ð-°°A±°s¸CÐ/@ÓAÑA€Kà”8˜Q ™U C¨Ð-°°A±°s¸CÐ/@ÓAÑA€KØÐr,   c                ó   — t        | |||«      S )z(Count flops for self-attention backward.©ræ   )r€   rŠ   r‹   rŒ   r1   r2   r3   s          r!   Úsdpa_backward_flopré   	  s   € ô
 $ N°KÀÈKÓXÐXr,   c
           
      óL   — t        |||| ||||	¬«      }t        d„ |D «       «      S )N)r³   r´   rµ   r«   r¶   r·   r¸   r¹   c              3   óB   K  — | ]  \  }}}}t        ||||«      –— Œ y ­wr   rè   ©r   rŠ   r‹   rŒ   r€   s        r!   r"   z1_flash_attention_backward_flop.<locals>.<genexpr>+  ó,   è ø€ ð áCIÑ?ˆK˜ K°ô 	! °¸iÈ×UÙCIùó   ‚rØ   )r«   r³   r´   rµ   ÚoutÚ	logsumexpr¶   r·   r¸   r¹   r2   r3   Úshapess                r!   Ú_flash_attention_backward_floprò     sB   € ô" 3ØØØØØØØØô	€Fô ñ áCIóó ð r,   c
           
      óL   — t        |||| ||||	¬«      }t        d„ |D «       «      S )N)r³   r´   rµ   r«   rÊ   rË   rÌ   rÍ   c              3   óB   K  — | ]  \  }}}}t        ||||«      –— Œ y ­wr   rè   rì   s        r!   r"   z5_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                 ó,   — t        | t        «      s| fS | S r   )r'   Útuple)Úxs    r!   Únormalize_tuplerù   j  s   € Ü�aœÔØˆtˆØ€Hr,   )Ú ÚKÚMÚBÚTc                 ó�   — t        dt        t        t        «      dz
  t        t	        | «      «      dz
  dz  «      «      }t        |   S )Nr   r   rM   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   — t         j                  |«      }| d|z  z  d›}|t         |   z   S )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                 ó   — |dk(  ry| |z  d›S )Nr   ú0%z.2%© )ÚnumÚdenoms     r!   Úconvert_to_percent_strr  �  s   € Ø�‚zØØ�E‰k˜#ÐÐr,   c                 ó.   ‡ — t        ‰ «      ˆ fd„«       }|S )Nc                 óB   •— t        | «      \  }} ‰|Ž }t        ||«      S r   )r   r   )r2   Ú	flat_argsÚspecrï   r4   s       €r!   r5   z)_pytreeify_preserve_structure.<locals>.nf‡  s'   ø€ ä& tÓ,‰ˆ	�4Ù�ˆmˆÜ˜c 4Ó(Ð(r,   r   r6   s   ` r!   Ú_pytreeify_preserve_structurer  †  s    ø€ Ü
ˆ1ƒXó)ó ð)ð
 €Ir,   c                   óú   ‡ — e Zd ZdZ	 	 	 	 ddej
                  j                  eej
                  j                     z  dz  dede	de
eef   dz  ddf
ˆ fd„Zdefd	„Zde
ee
eef   f   fd
„Zdd„Zd„ Zd„ Zd„ Zˆ xZS )r   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_mappingr8   c                 ód  •— t         ‰| �  «        t        d„ «      | _        || _        || _        d | _        |€i }|�t        j                  dd¬«       i t        ¥|j                  «       D ��ci c]   \  }}|t        |dd«      r|n
t        |«      “Œ" c}}¥| _	        t        «       | _        y c c}}w )Nc                  ó    — t        t        «      S r   )r   Úintr  r,   r!   Ú<lambda>z*FlopCounterMode.__init__.<locals>.<lambda>«  s
   € Ì+ÔVYÔJZr,   z<mods argument is not needed anymore, you can stop passing itrM   )Ú
stacklevelÚ_get_rawF)ÚsuperÚ__init__r   Úflop_countsr  r  ÚmodeÚwarningsÚwarnr-   Úitemsr   r7   r   Úmod_tracker)Úselfr  r  r  r  rS   ÚvÚ	__class__s          €r!   r!  zFlopCounterMode.__init__¤  s¶   ø€ ô 	‰ÑÔÜ6AÑBZÓ6[ˆÔØˆŒ
ØˆŒØ-1ˆŒ	ØÐ!ØˆNØÐÜ�M‰MÐXÐefÕgð
Üð
àWe×WkÑWkÔWmÔnÑWmÉtÈqÐRSˆq”w˜q *¨eÔ4‘!¼-ÈÓ:JÑJÐWmÒnð
ˆÔô )›?ˆÕùó os   Á-%B,c                 óN   — t        | j                  d   j                  «       «      S )NÚGlobal)rÙ   r"  Úvalues©r(  s    r!   Úget_total_flopszFlopCounterMode.get_total_flops¹  s!   € Ü�4×#Ñ# HÑ-×4Ñ4Ó6Ó7Ð7r,   c                 ó|   — | j                   j                  «       D ��ci c]  \  }}|t        |«      “Œ c}}S c c}}w )a  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)r(  rS   r)  s      r!   Úget_flop_countszFlopCounterMode.get_flop_counts¼  s9   € ð (,×'7Ñ'7×'=Ñ'=Ô'?Ô@Ñ'?™t˜q !�”4˜“7‘
Ð'?Ò@Ð@ùÓ@s   ž8c                 ó  ‡ ‡
‡‡— |€‰ j                   }|€d}dd l}d|_        g d¢}g }‰ j                  «       Š
t	        ‰
«      ŠdŠˆ
ˆˆˆ fd„}t        ‰ j                  j                  «       «      D ]?  }|dk(  rŒ	|j                  d«      d	z   }||kD  rŒ# |||d	z
  «      }|j                  |«       ŒA d‰ j                  v r ‰s|D ]  }	d
|	d   z   |	d<   Œ  |dd«      |z   }t        |«      dk(  rg d¢g}|j                  ||d¬«      S )Ni?B r   T)ÚModuleÚFLOPz% TotalFc           	      ó€  •— t        ‰
j                  |    j                  «       «      }‰	|‰k\  z  Š	d|z  }g }|j                  || z   t	        |‰«      t        |‰«      g«       ‰
j                  |    j                  «       D ]<  \  }}|j                  |dz   t        |«      z   t	        |‰«      t        |‰«      g«       Œ> |S )NÚ z - )rÙ   r"  r-  Úappendr	  r  r&  r  )Úmod_namer  r˜   Úpaddingr-  rS   r)  Úglobal_flopsÚglobal_suffixÚis_global_subsumedr(  s          €€€€r!   Úprocess_modz.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Ö:‘��1Ø—‘Ø˜e‘O¤c¨!£fÑ,Ü+¨A¨}Ó=Ü*¨1¨lÓ;ðõ ð ;ð ˆMr,   r,  Ú.r   r7  )r,  Ú0r  )ÚleftÚrightrB  )ÚheadersÚcolalign)r  ÚtabulateÚPRESERVE_WHITESPACEr/  r  Úsortedr"  ÚkeysÚcountÚextendr±   )r(  r  rE  Úheaderr-  r>  ÚmodÚ	mod_depthÚ
cur_valuesrµ   r;  r<  r=  s   `         @@@r!   Ú	get_tablezFlopCounterMode.get_tableÈ  s&  û€ Øˆ=Ø—J‘JˆEØˆ=ØˆEó 	à'+ˆÔ$Ú.ˆØˆØ×+Ñ+Ó-ˆÜ& |Ó4ˆØ"Ð÷	ô, ˜$×*Ñ*×/Ñ/Ó1Ö2ˆCØ�hŠØØŸ	™	 #›¨Ñ*ˆIØ˜5Ò Øá$ S¨)°a©-Ó8ˆJØ�M‰M˜*Õ%ð 3ð �t×'Ñ'Ñ'Ñ0BÛ�Ø  q¡™>��a’ð  ñ ! ¨1Ó-°Ñ6ˆFäˆv‹;˜!ÒÚ+Ð,ˆFà× Ñ  °ÐB\Ð Ó]Ð]r,   c                 óÂ   — | j                   j                  «        | j                  j                  «        t	        | «      | _        | j
                  j                  «        | S r   )r"  Úclearr'  Ú	__enter__Ú_FlopCounterModer#  r.  s    r!   rR  zFlopCounterMode.__enter__  sG   € Ø×Ñ×ÑÔ Ø×Ñ×"Ñ"Ô$Ü$ TÓ*ˆŒ	Ø�	‰	×ÑÔØˆr,   c                 ó  — | j                   €t        d«      ‚ | j                   j                  |Ž }d | _         | j                  j                  «        | j                  r$t        | j                  | j                  «      «       |S )Nz<Internal error: FlopCounter.__exit__ called but mode is None)r#  rO   Ú__exit__r'  r  ÚprintrO  r  )r(  r2   r]   s      r!   rU  zFlopCounterMode.__exit__  sh   € Ø�9‰9ÐÜ Ð!_Ó`Ð`ØˆD�I‰I×Ñ Ð%ˆØˆŒ	Ø×Ñ×!Ñ!Ô#Ø�<Š<Ü�$—.‘. §¡Ó,Ô-Øˆr,   c                 óÔ   — || j                   v rY| j                   |   } ||i |¤d|i¤Ž}t        | j                  j                  «      D ]  }| j                  |   |xx   |z  cc<   Œ |S )Nr/   )r-   Úsetr'  Úparentsr"  )r(  Úfunc_packetrï   r2   r3   Úflop_count_funcr„   Úpars           r!   Ú_count_flopszFlopCounterMode._count_flops  so   € Ø˜$×,Ñ,Ñ,Ø"×0Ñ0°Ñ=ˆOÙ(¨$ÐF°&ÑFÀ#ÒFˆJÜ˜4×+Ñ+×3Ñ3Ö4�Ø× Ñ  Ñ% kÓ2°jÑ@Ô2ð 5àˆ
r,   )NrM   TNr   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   Únnr4  r~   r  Úboolr1  r	   r!  r/  r  r2  rO  rR  rU  r]  Ú__classcell__)r*  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                   ó4   — e Zd ZdZdeddfd„Zd„ Zd„ Zd	d„Zy)
rS  TÚcounterr8   Nc                 ó   — || _         y r   )rf  )r(  rf  s     r!   r!  z_FlopCounterMode.__init__#  s	   € Øˆ�r,   c                 ó   — ddl }|j                  | j                  j                  «      }| 5   ||Ž }ddd«       |j                  | j                  j                  «      }|| j                  _        |fS # 1 sw Y   ŒCxY w)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)Úcopyrf  r"  )r(  Ú	branch_fnÚoperandsri  Úcheckpointed_flop_countsÚresultr"  s          r!   Ú$_execute_with_isolated_flop_countingz5_FlopCounterMode._execute_with_isolated_flop_counting&  si   € ó 	Ø#'§9¡9¨T¯\©\×-EÑ-EÓ#FÐ ÚÙ Ð)ˆF÷ à—i‘i §¡× 8Ñ 8Ó9ˆØ#;ˆ�‰Ô Ø�{Ð"Ð"÷	 ˆTús   ¬A4Á4A=c                 óV  — |t         j                  j                  j                  t         j                  j                  j                  hv }|rhddlm} ddlm}  ||d   «      }t        ||«      s't        |d«      r|j                  }nnt        ||«      sŒ'| j                  j                  |d ||«      S |t         j                  j                  j                  u �rI|\  }	}
}}| j                  |
|«      \  }}|t         u rt         S | j                  ||«      \  }}|t         u rt         S t#        |j%                  «       «      t#        |j%                  «       «      z  }i }|D ]€  }||   }||   }i }t#        |j%                  «       «      t#        |j%                  «       «      z  }|D ]5  }|j'                  |d«      }|j'                  |d«      }t)        ||«      ||<   Œ7 |||<   Œ‚ |j+                  «       D ]-  \  }}| j                  j,                  |   j/                  |«       Œ/ |S t         S )Nr   )Ú
get_kernelr   Ú
kernel_idxÚfn)r   ÚopsÚhigher_orderÚtriton_kernel_wrapper_mutationÚ triton_kernel_wrapper_functionalÚ*torch._higher_order_ops.triton_kernel_wraprp  Útriton.runtime.jitr   r'   Úhasattrrr  rf  r]  Úcondrn  ÚNotImplementedrX  rH  Úgetr   r&  r"  Úupdate)r(  ÚfuncÚtypesr2   r3   Ú	is_tritonrp  r   Úkernel_nameÚpredÚtrue_branchÚfalse_branchrk  Ú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_opsz)_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Ô:Ü˜;¨Ô-Ø"-§.¡.‘Kàô	 ! ¨kÕ:ð
 —<‘<×,Ñ,¨[¸$ÀÀfÓMÐMØ”U—Y‘Y×+Ñ+×0Ñ0Ò0ð
 9=Ñ5ˆD�+˜|¨Xà)-×)RÑ)RØ˜Xó*Ñ&ˆHÐ&ð œ>Ñ)Ü%Ð%à+/×+TÑ+TØ˜hó,Ñ(ˆ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                 óB  — |r|ni }|t         j                  j                  j                  j                  t         j                  j                  j
                  j                  t         j                  j                  j
                  j                  t         j                  j                  j                  j                  t         j                  j                  j                  j                  t         j                  j                  j                  j                  t         j                  j                  j                  j                  t         j                  j                  j                  j                  t         j                  j                  j                  j                  t         j                  j                  j                  j                  t         j                  j                  j                  j                  t         j                  j                  j                  j                  t         j                  j                  j                   j                  t         j                  j                  j"                  j                  t         j                  j$                  j&                  j                  hv rt(        S t+        |t         j,                  j.                  «      r| j1                  ||||«      S || j2                  j4                  vra|t         j                  j$                  j6                  j                  ur1| 5   |j8                  |i |¤Ž}|t(        ur|cd d d «       S 	 d d d «        ||i |¤Ž}| j2                  j;                  |j<                  |||«      S # 1 sw Y   Œ9xY wr   )r   rs  Ú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”  rf  r-   r¤   Ú	decomposer]  Ú_overloadpacket)r(  r~  r  r2   r3   Úrrï   s          r!   Ú__torch_dispatch__z#_FlopCounterMode.__torch_dispatch__w  s@  € Ù!‘ 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°°u¸dÀFÓKÐKð �t—|‘|×1Ñ1Ñ1°dÄ%Ç)Á)Ç.Á.×BWÑBW×B_ÑB_Ñ6_ÚØ"�D—N‘N DÐ3¨FÑ3�ØœNÑ*Ø÷ ‘à*÷ ñ �DÐ#˜FÑ#ˆØ�|‰|×(Ñ(¨×)=Ñ)=¸sÀDÈ&ÓQÐQ÷ �ús   Ì6NÎN)r  N)	r^  r_  r`  Úsupports_higher_order_operatorsr   r!  rn  r”  r«  r  r,   r!   rS  rS     s,   „ Ø&*Ð#ð ð °Dó ò#ò(;"ôz"Rr,   rS  )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   Ú	getLoggerr^  Úlogrx  r   r?   ÚImportErrorÚanyÚwarningrs  r–  r+   r-   r1  Ú__annotations__r7   r   Úmmr  rV   Úaddmmr[   Úbmmr`   Úbaddbmmrb   Ú
_scaled_mmrj   r~   rc  rt   Ú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   rS  r  r,   r!   Ú<module>rÓ     s>  ðæ Û Û ß FÑ FÝ )ß Ý $Ý $Ý 'Ý #Ý :Ý Ý Û àÐ5Ð
6€áˆTƒ]€Ùˆtƒ_€à€g×Ñ˜Ó!€ðÝ>ð ‡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 Ñ'Q:Ñ9Q: