§
    ŠŠtjÊ,  ã                   ó
  — d dl Z d dlmZ d dlZd dlmZ d dlmZmZm	Z	m
Z
 d dlmZ d„ Zdd„Zd	„ Zdd
„Zd„ Ze j        ded         fd„¦   «         Ze j        ded         fd„¦   «         Ze j        ded         fd„¦   «         ZdS )é    N)Ú	Generator)Úglobal_decomposition_table)Ú_rnn_helperÚgather_paramsÚgru_cellÚ	lstm_cell)Ú
while_loopc                 óî   ‡— g }g }| D ]_Š‰�/t          ˆfd„|D ¦   «         ¦  «        r‰                     ¦   «         Š|                     ‰¦  «         ‰�|                     ‰¦  «         Œ`t          |¦  «        S )Nc              3   óX   •K  — | ]$}t           j                             ‰|¦  «        V — Œ%d S )N)ÚtorchÚ_CÚ	_overlaps)Ú.0Úseen_tensorÚtensors     €úS/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/torch/export/_patches.pyú	<genexpr>z,_clone_tensors_that_alias.<locals>.<genexpr>   sF   øè è € ð &
ð &
Ø8C�EŒH×Ò˜v {Ñ3Ô3ð&
ð &
ð &
ð &
ð &
ð &
ó    )ÚanyÚcloneÚappendÚtuple)ÚtensorsÚcloned_tensorsÚseen_tensorsr   s      @r   Ú_clone_tensors_that_aliasr   
   s¥   ø€ ð €NØ€LØð (ð (ˆØÐ¥#ð &
ð &
ð &
ð &
ØGSð&
ñ &
ô &
ñ #
ô #
Ðð —\’\‘^”^ˆFØ×Ò˜fÑ%Ô%Ð%ØÐØ×Ò Ñ'Ô'Ð'øÝ�Ñ Ô Ð r   Fc                 óŽ  ‡‡‡‡— |d         }|d         Š|r|d         nd}|r|d         ndŠt          |¦  «        dk    r|d         nt          |¦  «        dk    r|d         ndŠ|d                              d¦  «        }|d                              d¦  «        }t          j        j                             | ||¦  «        Š|r‰                     d¦  «        n‰Št          ‰‰‰f¦  «        \  ŠŠŠt          j        ‰ 	                    d¦  «        gt          |j        dd…         ¦  «        ¢R |j        |j        dœŽ}	ˆfd	„}
ˆˆˆˆfd
„}t          j        dt          j        ¬¦  «        }t!          |
|||	||g¦  «        \  }}}}|r|                     d¦  «        }||                     d¦  «        |                     d¦  «        ffS )ay  
    1 layer fn for while loop LSTM

    Args:
        inp: Input tensor of shape (seq_len, batch, input_size)
        hidden: Tuple of (hx, cx) hidden states
        params: List of weight and bias tensors
        has_biases: Whether biases are included
        reverse: Whether to process sequence in reverse

    Returns:
        Tuple of (output, (final_hx, final_cx))
    r   é   é   Né   é   é   ©ÚdtypeÚdevicec                 ó6   •— | ‰                      d¦  «        k     S ©Nr   ©Úsize)ÚiÚoutÚhxÚcxÚprecomputed_inputs       €r   Úcond_fnz*one_layer_while_loop_lstm.<locals>.cond_fnA   ó   ø€ ØÐ$×)Ò)¨!Ñ,Ô,Ò,Ð,r   c           	      óT  •— |                       ¦   «         }t          j        |¦  «         t          j        |‰                     d¦  «        dz
  ¬¦  «         t	          ‰|         ||‰‰‰d¬¦  «        \  }}|                     ¦   «         }|                     d¦  «        ||<   | dz   |||fS )Nr   r   ©Úmaxr   )Ú	chunk_dim)Úitemr   Ú_check_is_sizer)   r   r   Úsqueeze)	Úidxr+   r,   r-   r*   Úhh_biasÚ	hh_weightÚ	hr_weightr.   s	        €€€€r   Úbody_fnz*one_layer_while_loop_lstm.<locals>.body_fnD   s­   ø€ à�HŠH‰JŒJˆÝÔ˜QÑÔÐÝÔ˜QÐ$5×$:Ò$:¸1Ñ$=Ô$=ÀÑ$AÐBÑBÔBÐBÝØ˜aÔ  " b¨)°W¸iÐSTð
ñ 
ô 
‰ˆˆBð �iŠi‰kŒkˆà—’˜A‘”ˆˆA‰Ø�Q‰w˜˜R Ð#Ð#r   ©r$   )ÚlenÚ	unsqueezer   ÚnnÚ
functionalÚlinearÚflipr   Úemptyr)   r   Úshaper$   r%   r   Úint64r	   r7   )ÚinpÚhiddenÚparamsÚ
has_biasesÚreverseÚ	ih_weightÚih_biasr,   r-   Ústep_outputr/   r<   ÚcntÚ_r+   Úfinal_hxÚfinal_cxr9   r:   r;   r.   s                    @@@@r   Úone_layer_while_loop_lstmrS      s"  øøøø€ ð �q”	€IØ�q”	€IØ%Ð/ˆf�QŒiˆi¨4€GØ%Ð/ˆf�QŒiˆi¨4€Gå˜‘[”[ AÒ%Ð%ˆˆqŒ	ˆ	½¸F¹¼ÀqÒ8HÐ8H¨6°!¬9¨9Èdð ð 
�Œ×	Ò	˜QÑ	Ô	€BØ	�Œ×	Ò	˜QÑ	Ô	€BåœÔ+×2Ò2°3¸	À7ÑKÔKÐØ5<ÐSÐ)×.Ò.¨qÑ1Ô1Ð1ÐBSÐÝ$=Ø	�G˜YÐ'ñ%ô %Ñ!€Iˆw˜	õ
 ”+Ø×Ò˜qÑ!Ô!ðå	ˆrŒx˜˜˜Œ|Ñ	Ô	ðð ð ŒhØŒyð	ð ð €Kð-ð -ð -ð -ð -ð$ð $ð $ð $ð $ð $ð $ð $õ Œ,�q¥¤Ð
,Ñ
,Ô
,€CÝ!+Ø�˜3 ¨R°Ð4ñ"ô "Ñ€A€sˆH�hð ð Ø�hŠh�q‰kŒkˆà�×!Ò! !Ñ$Ô$ h×&6Ò&6°qÑ&9Ô&9Ð:Ð:Ð:r   c	                 ó  — t          |¦  «        dk    rt          d¦  «        ‚t          |||d                              d¦  «        |d                              d¦  «        k    ¦  «        }t	          t          |d         |d         ¦  «        ¦  «        }	t          }
t          | |	||||||||
¦
  «
        \  }}t	          t          |Ž ¦  «        }|t          j	        |d         d¦  «        t          j	        |d         d¦  «        fS )aª  
    LSTM implementation using while_loop for export compatibility.

    This is a drop-in replacement for the default LSTM decomposition that uses
    while_loop instead of Python loops, making it more suitable for torch.export.

    Args:
        input: Input tensor
        hx: Tuple of (h0, c0) hidden states
        params: List of weight and bias tensors
        has_biases: Whether biases are included
        num_layers: Number of LSTM layers
        dropout: Dropout probability
        train: Training mode
        bidirectional: Whether to use bidirectional LSTM
        batch_first: Whether batch dimension is first

    Returns:
        Tuple of (output, h_n, c_n)
    r   zlstm expects two hidden statesr   r   )
r>   ÚAssertionErrorr   r)   ÚlistÚziprS   r   r   Ústack©Úinputr,   rI   rJ   Ú
num_layersÚdropoutÚtrainÚbidirectionalÚbatch_firstrH   Úlayer_fnr+   Úfinal_hiddenss                r   Úlstm_while_loop_implrb   [   sî   € õ> ˆ2�w„w�!‚|€|ÝÐ=Ñ>Ô>Ð>Ý˜6 :¨r°!¬u¯zªz¸!©}¬}ÀÀ1ÄÇ
Â
È1ÁÄÒ/MÑNÔN€FÝ•#�b˜”e˜R œUÑ#Ô#Ñ$Ô$€FÝ(€HÝ$ØØØØØØØØØØñô Ñ€Cˆõ �˜mÐ,Ñ-Ô-€MØ•”˜M¨!Ô,¨aÑ0Ô0µ%´+¸mÈAÔ>NÐPQÑ2RÔ2RÐRÐRr   c                 ó¦  ‡‡‡— |d         }|d         Š|r|d         nd}|r|d         ndŠt           j        j                             | ||¦  «        Š|r‰                     d¦  «        n‰Št          ‰‰f¦  «        \  ŠŠ|                     d¦  «        }t          j        ‰                     d¦  «        gt          |j
        dd…         ¦  «        ¢R |j        |j        dœŽ}ˆfd„}	ˆˆˆfd„}
t          j        dt           j        ¬	¦  «        }t          |	|
|||g¦  «        \  }}}|r|                     d¦  «        }||                     d¦  «        fS )
ad  
    1 layer fn for while loop GRU

    Args:
        inp: Input tensor of shape (seq_len, batch, input_size)
        hidden: Hidden state tensor
        params: List of weight and bias tensors
        has_biases: Whether biases are included
        reverse: Whether to process sequence in reverse

    Returns:
        Tuple of (output, final_hidden)
    r   r   r   Nr    r#   c                 ó6   •— | ‰                      d¦  «        k     S r'   r(   )r*   r+   Ú
cur_hiddenr.   s      €r   r/   z)one_layer_while_loop_gru.<locals>.cond_fn¯   r0   r   c                 óH  •— |                       ¦   «         }t          j        |¦  «         t          j        |‰                     d¦  «        dz
  ¬¦  «         t	          ‰|         |d d ‰‰¦  «        }|                     ¦   «         }|                     d¦  «        ||<   | dz   ||fS )Nr   r   r2   )r5   r   r6   r)   r   r   r7   )r8   r+   re   r*   r9   r:   r.   s       €€€r   r<   z)one_layer_while_loop_gru.<locals>.body_fn²   s£   ø€ à�HŠH‰JŒJˆÝÔ˜QÑÔÐÝÔ˜QÐ$5×$:Ò$:¸1Ñ$=Ô$=ÀÑ$AÐBÑBÔBÐBÝØ˜aÔ  *¨d°D¸)ÀWñ
ô 
ˆ
ð �iŠi‰kŒkˆØ×#Ò# AÑ&Ô&ˆˆA‰Ø�Q‰w˜˜ZÐ'Ð'r   r=   )r   r@   rA   rB   rC   r   r?   rD   r)   r   rE   r$   r%   r   rF   r	   r7   )rG   rH   rI   rJ   rK   rL   rM   re   rN   r/   r<   rO   rP   r+   Úfinal_hiddenr9   r:   r.   s                  @@@r   Úone_layer_while_loop_grurh   �   s¤  øøø€ ð �q”	€IØ�q”	€IØ%Ð/ˆf�QŒiˆi¨4€GØ%Ð/ˆf�QŒiˆi¨4€GåœÔ+×2Ò2°3¸	À7ÑKÔKÐØ5<ÐSÐ)×.Ò.¨qÑ1Ô1Ð1ÐBSÐÝ2°I¸wÐ3GÑHÔHÑ€IˆwØ×!Ò! !Ñ$Ô$€Jõ ”+Ø×Ò˜qÑ!Ô!ðå	ˆzÔ   Ô#Ñ	$Ô	$ðð ð ÔØÔ ð	ð ð €Kð-ð -ð -ð -ð -ð
(ð 
(ð 
(ð 
(ð 
(ð 
(ð 
(õ Œ,�q¥¤Ð
,Ñ
,Ô
,€CÝ% g¨w¸¸kÈ:Ð8VÑWÔWÑ€A€sˆLØð Ø�hŠh�q‰kŒkˆà�×$Ò$ QÑ'Ô'Ð'Ð'r   c	                 óÚ   — t          ||d¦  «        }t          |                     d¦  «        ¦  «        }	t          }
t	          | |	||||||||
¦
  «
        \  }}|t          j        |d¦  «        fS )a•  
    GRU implementation using while_loop for export compatibility.

    This is a drop-in replacement for the default GRU decomposition that uses
    while_loop instead of Python loops, making it more suitable for torch.export.

    Args:
        input: Input tensor
        hx: Hidden state tensor
        params: List of weight and bias tensors
        has_biases: Whether biases are included
        num_layers: Number of GRU layers
        dropout: Dropout probability
        train: Training mode
        bidirectional: Whether to use bidirectional GRU
        batch_first: Whether batch dimension is first

    Returns:
        Tuple of (output, h_n)
    Fr   )r   rV   Úunbindrh   r   r   rX   rY   s                r   Úgru_while_loop_implrk   Æ   s|   € õ> ˜6 :¨uÑ5Ô5€FÝ�"—)’)˜A‘,”,ÑÔ€FÝ'€HÝ$ØØØØØØØØØØñô Ñ€Cˆð •”˜M¨1Ñ-Ô-Ð-Ð-r   Úreturn)NNNc              #   óÂ  K  — t           d         }|                     | d¦  «        }| j                             t          j        j        j        d¦  «        }	 ||| <   || j        t          j        j        j        <   dV — |�||| <   n|                     | d¦  «         |� || j        t          j        j        j        <   dS | j                             t          j        j        j        d¦  «         dS # |�||| <   n|                     | d¦  «         |�|| j        t          j        j        j        <   w | j                             t          j        j        j        d¦  «         w xY w)a†  
    Generic context manager for registering while_loop-based RNN decompositions.

    Args:
        rnn_op: The aten operation to patch (e.g., torch.ops.aten.lstm.input)
        rnn_impl: The while_loop-based implementation function

    Note:
        This is an internal helper. Use register_lstm_while_loop_decomposition()
        or register_gru_while_loop_decomposition() instead.
    Úpost_autogradN)r   ÚgetÚ
py_kernelsr   r   ÚDispatchKeyÚCompositeImplicitAutogradÚpop)Úrnn_opÚrnn_implÚregistryÚoriginal_decompÚoriginal_py_kernels        r   Ú&_register_rnn_while_loop_decompositionry   ÷   sq  è è € õ *¨/Ô:€Hð —l’l 6¨4Ñ0Ô0€Oð  Ô*×.Ò.ÝŒÔÔ6¸ñô ÐðXà#ˆ�ÑØLTˆÔ�%œ(Ô.ÔHÑIØˆˆˆð Ð&Ø.ˆH�VÑÐð �LŠL˜ Ñ&Ô&Ð&ð Ð)à"ð Ô�eœhÔ2ÔLÑMÐMÐMð
 Ô×!Ò!¥%¤(Ô"6Ô"PÐRVÑWÔWÐWÐWÐWøð Ð&Ø.ˆH�VÑÐð �LŠL˜ Ñ&Ô&Ð&ð Ð)à"ð Ô�eœhÔ2ÔLÑMÐMð
 Ô×!Ò!¥%¤(Ô"6Ô"PÐRVÑWÔWÐWÐWøøøs   Á'C. Ã.A0Ec               #   ó    K  — t          t          j        j        j        j        t          ¦  «        5  dV — ddd¦  «         dS # 1 swxY w Y   dS )aŽ  
    Context manager that temporarily registers the while_loop-based LSTM decomposition.

    The while_loop-based decomposition is more suitable for export and graph-based
    execution, as it avoids Python control flow that cannot be captured in the graph.
    This should support dynamic sequence lengths, however as while_loop does not
    support Autograd yet, an ExportedProgram created with this will not be trainable.

    Usage::

        from torch.export._patches import register_lstm_while_loop_decomposition
        from torch.export import export

        with register_lstm_while_loop_decomposition():
            # Export your model with LSTM
            ep = export(model, (x, h0, c0))

    Note:
        This context manager temporarily modifies the global decomposition table
        and py_kernels registration. The original registrations are restored when
        exiting the context.
    N)ry   r   ÚopsÚatenÚlstmrZ   rb   © r   r   Ú&register_lstm_while_loop_decompositionr   '  s—   è è € õ0 
0ÝŒ	ŒÔÔ!Õ#7ñ
ô 
ð ð ð 	ˆˆˆðð ð ñ ô ð ð ð ð ð ð ð øøøð ð ð ð ð ð ó   ±AÁAÁ
Ac               #   ó    K  — t          t          j        j        j        j        t          ¦  «        5  dV — ddd¦  «         dS # 1 swxY w Y   dS )a†  
    Context manager that temporarily registers the while_loop-based GRU decomposition.

    The while_loop-based decomposition is more suitable for export and graph-based
    execution, as it avoids Python control flow that cannot be captured in the graph.
    This should support dynamic sequence lengths, however as while_loop does not
    support Autograd yet, an ExportedProgram created with this will not be trainable.

    Usage::

        from torch.export._patches import register_gru_while_loop_decomposition
        from torch.export import export

        with register_gru_while_loop_decomposition():
            # Export your model with GRU
            ep = export(model, (x, h0))

    Note:
        This context manager temporarily modifies the global decomposition table
        and py_kernels registration. The original registrations are restored when
        exiting the context.
    N)ry   r   r{   r|   ÚgrurZ   rk   r~   r   r   Ú%register_gru_while_loop_decompositionrƒ   E  s—   è è € õ0 
0ÝŒ	ŒÔÔ Õ"5ñ
ô 
ð ð ð 	ˆˆˆðð ð ñ ô ð ð ð ð ð ð ð øøøð ð ð ð ð ð r€   )F)Ú
contextlibÚcollections.abcr   r   Útorch._decompr   Útorch._decomp.decompositionsr   r   r   r   Ú"torch._higher_order_ops.while_loopr	   r   rS   rb   rh   rk   Úcontextmanagerry   r   rƒ   r~   r   r   ú<module>rŠ      st  ðØ Ð Ð Ð Ø %Ð %Ð %Ð %Ð %Ð %à €€€Ø 4Ð 4Ð 4Ð 4Ð 4Ð 4Ø XÐ XÐ XÐ XÐ XÐ XÐ XÐ XÐ XÐ XÐ XÐ XØ 9Ð 9Ð 9Ð 9Ð 9Ð 9ð!ð !ð !ð >;ð >;ð >;ð >;ðB1Sð 1Sð 1Sðh4(ð 4(ð 4(ð 4(ðn..ð ..ð ..ðb Ôð,XàÐÔ ð,Xð ,Xð ,Xñ Ôð,Xð^ Ôð°	Ð:JÔ0Kð ð ð ñ Ôðð: Ôð¨yÐ9IÔ/Jð ð ð ñ Ôðð ð r   