§
    ŠŠtj  ã                   ó²  — d dl mZ d dlZd dlmZ d dlmZ d dlmZ d dl	m
Z
mZmZ d dlmZmZ g d¢Z ed	¬
¦  «        dedededededdfd„¦   «         Z ed	¬
¦  «         G d„ de¦  «        ¦   «         Z ed	¬
¦  «        	 ddedeej                 dz  ddfd„¦   «         Z ed	¬
¦  «        dedefd„¦   «         Z ed	¬
¦  «        dededefd„¦   «         ZdS )é    )Ú
NamedTupleN)Úcompatibility)ÚGraph)ÚGraphModule)Úmap_argÚNodeÚTarget)Ú	ShapePropÚTensorMetadata)Ú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_targetÚreturnc                 ó,  ‡	— t          ¦   «         }i Š	| j        j        D ]î}|j        |k    rÅ|j        |k    rºt          |j        ˆ	fd„¦  «        }t          |j        ˆ	fd„¦  «        }t          |t          ¦  «        st          dt          |¦  «        › �¦  «        ‚t          |t          ¦  «        st          dt          |¦  «        › �¦  «        ‚|                     |||||j        ¦  «        ‰	|<   ŒÒ|                     |ˆ	fd„¦  «        ‰	|<   Œï|| _        dS )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 target.
    c                 ó   •— ‰|          S ©N© ©ÚnÚval_maps    €ú`/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/torch/fx/passes/graph_manipulation.pyú<lambda>z+replace_target_nodes_with.<locals>.<lambda>$   s   ø€ °¸´
€ ó    c                 ó   •— ‰|          S r   r   r   s    €r   r    z+replace_target_nodes_with.<locals>.<lambda>%   s   ø€ °G¸A´J€ r!   zExpected tuple, got zExpected dict, got c                 ó   •— ‰|          S r   r   r   s    €r   r    z+replace_target_nodes_with.<locals>.<lambda>.   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      s2  ø€ õ ‘”€IØ "€GØ”Ô%ð Lð LˆØŒ7�fÒÐ ¤°
Ò!:Ð!:Ý˜4œ9Ð&:Ð&:Ð&:Ð&:Ñ;Ô;ˆDÝ˜Tœ[Ð*>Ð*>Ð*>Ð*>Ñ?Ô?ˆFÝ˜d¥EÑ*Ô*ð JÝ$Ð%H½DÀ¹J¼JÐ%HÐ%HÑIÔIÐIÝ˜f¥dÑ+Ô+ð KÝ$Ð%I½4À¹<¼<Ð%IÐ%IÑJÔJÐJØ%×1Ò1Ø˜
 D¨&°$´)ñô ˆG�D‰MˆMð &×/Ò/°Ð6JÐ6JÐ6JÐ6JÑKÔKˆG�D‰MˆMØ€I„O€O€Or!   c                   ó$   — e Zd ZU eed<   eed<   dS )r   Úoutput_sizeÚ
total_sizeN)Ú__name__Ú
__module__Ú__qualname__ÚintÚ__annotations__r   r!   r   r   r   2   s%   € € € € € € àÐÐÑØ€O€O�O€O€Or!   r   r(   c                 óš   — |� t          | ¦  «        j        |Ž  | j        j        D ]$}|j        dk    r nt          | |¦  «        |_        Œ%dS )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(   r3   s      r   r   r   8   s`   € ð Ðà&�	�)ÑÔÔ&¨Ð-Ð-à”Ô%ð <ð <ˆØŒ7�hÒÐØˆEÝ*¨9°dÑ;Ô;ˆŒˆØ
€Fr!   r3   c                 ód   — | j                              d¦  «        }|st          d| › d�¦  «        ‚|S )NÚtensor_metazNode zQ has no tensor metadata associated with it! Check that shape propagation has run.)ÚmetaÚgetÚRuntimeError)r3   r@   s     r   r   r   J   sM   € à”)—-’- Ñ.Ô.€Kàð 
Ýð5�Dð 5ð 5ð 5ñ
ô 
ð 	
ð
 Ðr!   c                 ó0  — d}|j         dk    rat          |                      ¦   «         ¦  «        }||j                 }|                     ¦   «         }|D ]\  }}||                     ¦   «         z  }Œt          |¦  «        }|j                             ¦   «         }	||	z  }|j        r.t          j
        g |j        ¬¦  «                             ¦   «         }
n-t          j        g |j        ¬¦  «                             ¦   «         }
|
|z  }|
|	z  }t          ||¦  «        S )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_quantizedrF   Úelement_sizeÚtensorr   )r   r3   Útotal_num_of_elemsÚsubmodule_dictÚ	submoduleÚ
parametersÚ_nameÚpr@   Úoutput_elemÚsize_per_elem_bytesr6   r5   s                r   r   r   W   s)  € ð Ðà„w�-ÒÐÝ˜i×5Ò5Ñ7Ô7Ñ8Ô8ˆØ" 4¤;Ô/ˆ	Ø×/Ò/Ñ1Ô1ˆ
à"ð 	,ð 	,‰HˆE�1Ø !§'¢'¡)¤)Ñ+ÐÐõ " $Ñ'Ô'€KØÔ#×)Ò)Ñ+Ô+€KØ˜+Ñ%ÐàÔð WÝ#Ô;Ø�kÔ'ð
ñ 
ô 
ç
Š,‰.Œ.ð 	Ðõ $œl¨2°[Ô5FÐGÑGÔG×TÒTÑVÔVÐØ$Ð'9Ñ9€JØ%¨Ñ3€KÝ�k :Ñ.Ô.Ð.r!   r   )Útypingr   rL   Útorch.fx._compatibilityr   Útorch.fx.graphr   Útorch.fx.graph_moduler   Útorch.fx.noder   r   r	   Útorch.fx.passes.shape_propr
   r   Ú__all__Ústrr   r   ÚlistÚTensorr   r   r   r   r!   r   ú<module>rb      s*  ðØ Ð Ð Ð Ð Ð à €€€Ø 1Ð 1Ð 1Ð 1Ð 1Ð 1Ø  Ð  Ð  Ð  Ð  Ð  Ø -Ð -Ð -Ð -Ð -Ð -Ø /Ð /Ð /Ð /Ð /Ð /Ð /Ð /Ð /Ð /Ø @Ð @Ð @Ð @Ð @Ð @Ð @Ð @ðð ð €ð € eÐ,Ñ,Ô,ð Øð àð ð ð ð ð	 ð
 ð ð 
ð ð  ð  ñ -Ô,ð ð: € eÐ,Ñ,Ô,ðð ð ð ð �ñ ô ñ -Ô,ðð
 € eÐ,Ñ,Ô,à>Bðð ØðØ"& u¤|Ô"4°tÑ";ðà	ðð ð ñ -Ô,ðð" € eÐ,Ñ,Ô,ð	˜$ð 	 >ð 	ð 	ð 	ñ -Ô,ð	ð € eÐ,Ñ,Ô,ð/ ð /°4ð /¸Jð /ð /ð /ñ -Ô,ð/ð /ð /r!   