ó
    Eñi  ã                   ó^  • S SK JrJrJr  S SKrS SKJr  S SKJr  S SK	J
r
  S SKJrJrJr  S SKJr  / SQr\" S	S
9S\
S\S\S\S\4
S j5       r\" S	S
9 " S S\5      5       r\" S	S
9 SS\
S\\\R,                        SS4S jj5       r\" S	S
9S\S\4S j5       r\" S	S
9S\
S\S\4S j5       rg)é    )ÚAnyÚ
NamedTupleÚOptionalN)Úcompatibility)ÚGraph)ÚGraphModule)Úmap_argÚNodeÚTarget)Ú	ShapeProp)Úreplace_target_nodes_withÚ
size_bytesÚget_size_of_all_nodesÚget_tensor_metaÚget_size_of_nodeF)Úis_backward_compatibleÚ	fx_moduleÚold_opÚ
old_targetÚnew_opÚ
new_targetc                 ó2  ^	• [        5       n0 m	U R                  R                   Hê  nUR                  U:X  a¾  UR                  U:X  a®  [        UR                  U	4S j5      n[        UR                  U	4S j5      n[        U[        5      (       d  [        S[        U5       35      e[        U[        5      (       d  [        S[        U5       35      eUR                  X4XxUR                  5      T	U'   MÑ  UR                  UU	4S j5      T	U'   Mì     XPl        g)zŽModifies all nodes in fx_module.graph.nodes which match the specified op code and target,
and updates them to match the new op code and targetc                 ó   >• TU    $ ©N© ©ÚnÚval_maps    €Ú_/home/mande/repo/quber/.venv/lib/python3.13/site-packages/torch/fx/passes/graph_manipulation.pyÚ<lambda>Ú+replace_target_nodes_with.<locals>.<lambda>#   s	   ø€ °¸²
ó    c                 ó   >• TU    $ r   r   r   s    €r   r    r!   $   s	   ø€ °G¸A²Jr"   zExpected tuple, got zExpected dict, got c                 ó   >• TU    $ r   r   r   s    €r   r    r!   -   s	   ø€ ÀÈÂ
r"   N)r   ÚgraphÚnodesÚopÚtargetr	   ÚargsÚkwargsÚ
isinstanceÚtupleÚAssertionErrorÚtypeÚdictÚcreate_nodeÚnameÚ	node_copy)
r   r   r   r   r   Ú	new_graphÚnoder)   r*   r   s
            @r   r   r      så   ø€ ô “€IØ "€GØ—‘×%Ô%ˆØ�7‰7�fÓ §¡°
Ó!:Ü˜4Ÿ9™9Ô&:Ó;ˆDÜ˜TŸ[™[Ô*>Ó?ˆFÜ˜d¤E×*Ñ*Ü$Ð';¼DÀ»J¸<Ð%HÓIÐIÜ˜f¤d×+Ñ+Ü$Ð':¼4À»<¸.Ð%IÓJÐJØ%×1Ñ1Ø D°$·)±)óˆG�D‹Mð &×/Ñ/°Ô6JÓKˆG�D‹Mñ &ð  …Or"   c                   ó*   • \ rS rSr% \\S'   \\S'   Srg)r   é1   Úoutput_sizeÚ
total_sizer   N)Ú__name__Ú
__module__Ú__qualname__Ú__firstlineno__ÚintÚ__annotations__Ú__static_attributes__r   r"   r   r   r   1   s   ‡ àÓØ†Or"   r   r)   Úreturnc                 ó¸   • Ub  [        U 5      R                  " U6   U R                  R                   H%  nUR                  S:X  a    g[        X5      Ul        M'     g)zÀGiven a fx graph module, update each node with its total size (weights + bias + output)
and its output_size(output). For a non-module node, the total size is the output size.
return total sizeNÚoutput)r   Ú	propagater%   r&   r'   r   r   )r   r)   r4   s      r   r   r   7   sV   € ð Ñä�)Ó×&Ò&¨Ñ-à—‘×%Ô%ˆØ�7‰7�hÓØà
ô +¨9Ó;ˆŽñ &ð r"   r4   c                 óh   • U R                   R                  S5      nU(       d  [        SU  S35      eU$ )NÚtensor_metazNode zQ has no tensor metadata associated with it! Check that shape propagation has run.)ÚmetaÚgetÚRuntimeError)r4   rE   s     r   r   r   I   s=   € à—)‘)—-‘- Ó.€KæÜØ�D�6ð 4ð 5ó
ð 	
ð
 Ðr"   c                 ó  • SnUR                   S:X  aT  [        U R                  5       5      nX1R                     nUR	                  5       nU H  u  pgX'R                  5       -  nM     [        U5      nUR                  R                  5       n	X)-  nUR                  (       a.  [        R                  " / UR                  S9R                  5       n
O-[        R                  " / UR                  S9R                  5       n
X¢-  nX©-  n[        XË5      $ )z‚Given a node with node.dtype and node.shape, return its total size and its output size.
total_size = weights + bias + output_size
r   Úcall_module)Údtype)r'   r/   Únamed_modulesr(   Únamed_parametersÚnumelr   ÚshapeÚis_quantizedÚtorchÚ_empty_affine_quantizedrK   Úelement_sizeÚtensorr   )r   r4   Útotal_num_of_elemsÚsubmodule_dictÚ	submoduleÚ
parametersÚ_nameÚprE   Úoutput_elemÚsize_per_elem_bytesr8   r7   s                r   r   r   V   sò   € ð Ðà‡w�w�-ÓÜ˜i×5Ñ5Ó7Ó8ˆØ"§;¡;Ñ/ˆ	Ø×/Ñ/Ó1ˆ
ã"‰HˆEØ§'¡'£)Ñ+Òñ #ô " $Ó'€KØ×#Ñ#×)Ñ)Ó+€KØÑ%Ðà××Ü#×;Ò;Ø�k×'Ñ'ñ
ç
‰,‹.ñ 	ô $Ÿlšl¨2°[×5FÑ5FÑG×TÑTÓVÐØ$Ñ9€JØ%Ñ3€KÜ�kÓ.Ð.r"   r   )Útypingr   r   r   rQ   Útorch.fx._compatibilityr   Útorch.fx.graphr   Útorch.fx.graph_moduler   Útorch.fx.noder	   r
   r   Útorch.fx.passes.shape_propr   Ú__all__Ústrr   r   ÚlistÚTensorr   r   r   r   r"   r   Ú<module>rg      s%  ðç ,Ñ ,ã Ý 1Ý  Ý -ß /Ñ /Ý 0ò€ñ  eÑ,ð Øð àð ð ð ð ð	 ð
 ó ó -ð ñ6  eÑ,ô�ó ó -ðñ
  eÑ,àAEñØðØ"*¨4°·±Ñ+=Ñ">ðà	ôó -ðñ"  eÑ,ð	˜$ð 	 3ó 	ó -ð	ñ  eÑ,ð/ ð /°4ð /¸Jó /ó -ñ/r"   