§
    ŠŠtj”   ã                   ó8  — d dl Z d dlmZmZ d dl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 d dlmZ d dlmZmZ d	d
gZ ed¬¦  «         G d„ d	e¦  «        ¦   «         Z	 ddej        dedefd„Z ed¬¦  «         G d„ d
ej        j        ¦  «        ¦   «         ZdS )é    N)ÚAnyÚ
NamedTuple)Úenable_python_dispatcher)Údetect_fake_mode)Ú(is_contiguous_for_memory_format_or_false)Úis_sparse_any)Úcompatibility)Úmap_aggregateÚNodeÚTensorMetadataÚ	ShapePropT)Úis_backward_compatiblec                   óž   — e Zd ZU dZej        ed<   ej        ed<   eed<   e	e
df         ed<   ej        dz  ed<   eed	<   eeef         ed
<   dS )r   zUA structure containing pertinent information about a tensor within a PyTorch program.ÚshapeÚdtypeÚrequires_grad.ÚstrideNÚmemory_formatÚis_quantizedÚqparams)Ú__name__Ú
__module__Ú__qualname__Ú__doc__ÚtorchÚSizeÚ__annotations__r   ÚboolÚtupleÚintr   ÚdictÚstrr   © ó    úX/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/torch/fx/passes/shape_prop.pyr   r      sŒ   € € € € € € à_Ð_ð Œ:ÐÐÑØŒ;ÐÐÑØÐÐÑØ�#�s�(ŒOÐÐÑØÔ&¨Ñ-Ð-Ð-Ñ-ð ÐÐÑØ�#�s�(Œ^ÐÐÑÐÐr$   ÚresultÚinclude_contiguityÚreturnc           	      ó6  — | j         }| j        }| j        }t          | ¦  «        s|                      ¦   «         nd}d}|rLt          | ¦  «        s=t
          j        t
          j        t
          j        f}|D ]}t          | |¬¦  «        r|} nŒ| j
        }	i }
|	rð|                      ¦   «         }||
d<   |t
          j        t
          j        fv r/|                      ¦   «         |
d<   |                      ¦   «         |
d<   nŽ|t
          j        t
          j        t
          j        fv ri|                      ¦   «                              ¦   «         |
d<   |                      ¦   «                              ¦   «         |
d<   |                      ¦   «         |
d<   t/          ||||||	|
¦  «        S )zB
    Extract a TensorMetadata NamedTuple describing `result`.
    r#   N)r   ÚqschemeÚscaleÚ
zero_pointÚaxis)r   r   r   r   r   r   Úcontiguous_formatÚchannels_lastÚchannels_last_3dr   r   r*   Úper_tensor_affineÚper_tensor_symmetricÚq_scaleÚq_zero_pointÚper_channel_affineÚ per_channel_affine_float_qparamsÚper_channel_symmetricÚq_per_channel_scalesÚtolistÚq_per_channel_zero_pointsÚq_per_channel_axisr   )r&   r'   r   r   r   r   r   Úmemory_formatsÚquery_formatr   r   r*   s               r%   Ú_extract_tensor_metadatar>   %   s¾  € ð ŒL€EØŒL€EØÔ(€MÝ$1°&Ñ$9Ô$9ÐAˆV�]Š]‰_Œ_ˆ_¸r€Fà€Màð ¥-°Ñ"7Ô"7ð åÔ#ÝÔÝÔ"ð
ˆð
 +ð 	ð 	ˆLÝ7Ø lðñ ô ð ð !-�Ø�ð	ð Ô&€LØ €GØð :Ø—.’.Ñ"Ô"ˆØ$ˆ�	ÑØ•uÔ.µÔ0JÐKÐKÐKØ%Ÿ~š~Ñ/Ô/ˆG�GÑØ$*×$7Ò$7Ñ$9Ô$9ˆG�LÑ!Ð!ØÝÔ$ÝÔ2ÝÔ'ð
ð 
ð 
ð  &×:Ò:Ñ<Ô<×CÒCÑEÔEˆG�GÑØ$*×$DÒ$DÑ$FÔ$F×$MÒ$MÑ$OÔ$OˆG�LÑ!Ø$×7Ò7Ñ9Ô9ˆG�F‰OåØˆu�m V¨]¸LÈ'ñô ð r$   c                   ón   ‡ — e Zd ZdZddej        j        deddfˆ fd„Zde	defˆ fd„Z
d	edefˆ fd
„Zˆ xZS )r   aE  
    Execute an FX graph Node-by-Node and
    record the shape and type of the result
    into the corresponding node.

    Example:
         In this example, we record the shape
         and data type of a module given
         an example input ``torch.randn(50, D_in)``.
         We print the name, shape and dtype of each node.

        class TwoLayerNet(torch.nn.Module):
            def __init__(self, D_in, H, D_out):
                super().__init__()
                self.linear1 = torch.nn.Linear(D_in, H)
                self.linear2 = torch.nn.Linear(H, D_out)
            def forward(self, x):
                h_relu = self.linear1(x).clamp(min=0)
                y_pred = self.linear2(h_relu)
                return y_pred
        N, D_in, H, D_out = 64, 1000, 100, 10
        x = torch.randn(N, D_in)
        y = torch.randn(N, D_out)
        model = TwoLayerNet(D_in, H, D_out)
        gm = torch.fx.symbolic_trace(model)
        sample_input = torch.randn(50, D_in)
        ShapeProp(gm).propagate(sample_input)

        for node in gm.graph.nodes:
            print(node.name, node.meta['tensor_meta'].dtype,
                node.meta['tensor_meta'].shape)

        The output of this code is:

        x torch.float32 torch.Size([50, 1000])
        linear1 torch.float32 torch.Size([50, 100])
        clamp_1 torch.float32 torch.Size([50, 100])
        linear2 torch.float32 torch.Size([50, 10])
        output torch.float32 torch.Size([50, 10])

    Args:
         module (GraphModule): The module to be executed
         fake_mode (FakeTensorMode): A fake mode for copying the gm

    NÚgmÚ	fake_moder(   c                 óê   •— t          ¦   «                              |¦  «         |€t          ¦   «         }|�$ddlm}  || j        |¦  «        | _        || _        nd | _        d | _        | j        | _        d S )Nr   )Údeepcopy_to_fake_tensor)	ÚsuperÚ__init__r   Útorch._dynamo.utilsrC   ÚmoduleÚfake_modulerA   Úreal_module)Úselfr@   rA   rC   Ú	__class__s       €r%   rE   zShapeProp.__init__ˆ   s…   ø€ Ý‰Œ×Ò˜ÑÔÐØÐÝ(Ñ*Ô*ˆIØÐ ØCÐCÐCÐCÐCÐCð  7Ð6°t´{ÀIÑNÔNˆDÔØ&ˆDŒNˆNà#ˆDÔØ!ˆDŒNàœ;ˆÔÐÐr$   Únc                 ó^  •‡
— ddl m}m} 	 | j        �| j        | _        	 | j        �~| j        5  t          ¦   «         5  t          ¦   «                              |¦  «        } || j        j	        ||¦  «         d d d ¦  «         n# 1 swxY w Y   d d d ¦  «         n# 1 swxY w Y   n!t          ¦   «                              |¦  «        }| j
        | _        n# | j
        | _        w xY wnR# t          $ rE}t          j        ¦   «          t          d|                     ¦   «         › d|j        › �¦  «        |‚d }~ww xY wdŠ
dt"          dt"          fˆ
fd„}t%          ||¦  «        }‰
r
||j        d	<   | j        r&| j        j	        x}r |||¦  «        x}	r
|	|j        d
<   t'          |¦  «        |j        d<   |S )Nr   )Úcompute_unbacked_bindingsÚrebind_unbackedzShapeProp error for: node=z with meta=FÚobjr(   c                 ó^   •— t          | t          j        ¦  «        rdŠt          | ¦  «        S | S )NT)Ú
isinstancer   ÚTensorr>   )rP   Úfound_tensors    €r%   Úextract_tensor_metaz/ShapeProp.run_node.<locals>.extract_tensor_meta¼   s/   ø€ Ý˜#�uœ|Ñ,Ô,ð à#�Ý/°Ñ4Ô4Ð4à�
r$   Útensor_metaÚunbacked_bindingsÚtype)Ú%torch.fx.experimental.symbolic_shapesrN   rO   rH   rG   rA   r   rD   Úrun_nodeÚ	shape_envrI   Ú	ExceptionÚ	tracebackÚ	print_excÚRuntimeErrorÚformat_nodeÚmetar   r
   rX   )rJ   rL   rN   rO   r&   ÚerU   ra   r[   Úsymbol_to_pathrT   rK   s             @€r%   rZ   zShapeProp.run_node    sµ  øø€ ð	
ð 	
ð 	
ð 	
ð 	
ð 	
ð 	
ð 	
ð
	ØÔÐ+ð #Ô.�”ð/Ø”>Ð-Øœð Mð MÕ)AÑ)CÔ)Cð Mð MÝ!&¡¤×!1Ò!1°!Ñ!4Ô!4˜Ø'˜¨¬Ô(@À!ÀVÑLÔLÐLðMð Mð Mñ Mô Mð Mð Mð Mð Mð Mð Møøøð Mð Mð Mð Mð Mð Mð Mñ Mô Mð Mð Mð Mð Mð Mð Møøøð Mð Mð Mð Møõ #™WœW×-Ò-¨aÑ0Ô0�Fà"Ô.�”�ø˜dÔ.�”Ð.Ð.Ð.Ð.�øÝð 	ð 	ð 	ÝÔÑ!Ô!Ð!ÝØQ¨Q¯]ª]©_¬_ÐQÐQÈÌÐQÐQñô àðøøøøð	øøøð ˆð	¥Sð 	­Sð 	ð 	ð 	ð 	ð 	ð 	õ ˜VÐ%8Ñ9Ô9ˆØð 	)Ø$(ˆAŒF�=Ñ!àŒ>ð 	=Ø!œ^Ô5Ð5�	ð =Ø";Ð";¸IÀvÑ"NÔ"NÐN�ð=ð /=�”Ð*Ñ+å˜f™œˆŒˆv‰Øˆsu   ŒC%  C ®B½9BÁ6BÂB	ÂBÂ	B	Â
BÂC ÂBÂC Â BÂ!%C ÃC% ÃC!Ã!C% Ã%
D4Ã/A D/Ä/D4Úargsc                 ób   •‡ — ‰ j         �ˆ fd„|D ¦   «         }n|} t          ¦   «         j        |Ž S )a  
        Run `module` via interpretation and return the result and
        record the shape and type of each node.

        Args:
            *args (Tensor): the sample input.

        Returns:
            Any: The value returned from executing the Module
        Nc                 ó|   •— g | ]8}t          |t          j        ¦  «        r‰j                             |¦  «        n|‘Œ9S r#   )rR   r   rS   rA   Úfrom_tensor)Ú.0ÚtrJ   s     €r%   ú
<listcomp>z'ShapeProp.propagate.<locals>.<listcomp>Ý   sP   ø€ ð ð ð àõ 2<¸A½u¼|Ñ1LÔ1LÐS�”×*Ò*¨1Ñ-Ô-Ð-ÐRSðð ð r$   )rA   rD   Úrun)rJ   rd   Ú	fake_argsrK   s   `  €r%   Ú	propagatezShapeProp.propagateÑ   sR   øø€ ð Œ>Ð%ðð ð ð àðñ ô ˆIˆIð
 ˆIØ�u‰wŒwŒ{˜IÐ&Ð&r$   )N)r   r   r   r   r   ÚfxÚGraphModuler   rE   r   rZ   rm   Ú__classcell__)rK   s   @r%   r   r   X   sÀ   ø€ € € € € ð,ð ,ð\'ð '˜5œ8Ô/ð '¸Cð 'È4ð 'ð 'ð 'ð 'ð 'ð 'ð0/˜$ð / 3ð /ð /ð /ð /ð /ð /ðb'˜sð ' sð 'ð 'ð 'ð 'ð 'ð 'ð 'ð 'ð 'ð 'r$   )T)r]   Útypingr   r   r   Útorch.fxÚtorch._dispatch.pythonr   Útorch._guardsr   Útorch._prims_commonr   Útorch._subclasses.meta_utilsr   Útorch.fx._compatibilityr	   Útorch.fx.noder
   r   Ú__all__r   rS   r   r>   rn   ÚInterpreterr   r#   r$   r%   ú<module>r{      s“  ðØ Ð Ð Ð Ø "Ð "Ð "Ð "Ð "Ð "Ð "Ð "à €€€Ø €€€Ø ;Ð ;Ð ;Ð ;Ð ;Ð ;Ø *Ð *Ð *Ð *Ð *Ð *Ø HÐ HÐ HÐ HÐ HÐ HØ 6Ð 6Ð 6Ð 6Ð 6Ð 6Ø 1Ð 1Ð 1Ð 1Ð 1Ð 1Ø -Ð -Ð -Ð -Ð -Ð -Ð -Ð -ð ˜[Ð
)€ð € dÐ+Ñ+Ô+ðð ð ð ð �Zñ ô ñ ,Ô+ðð( 6:ð0ð 0ØŒLð0Ø.2ð0àð0ð 0ð 0ð 0ðf € dÐ+Ñ+Ô+ðJ'ð J'ð J'ð J'ð J'�”Ô$ñ J'ô J'ñ ,Ô+ðJ'ð J'ð J'r$   