Ë
    †\;jÔe  ã                   ó„  — d dl Z d dlZd dl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mZ d dlmZ ddlmZ d	„ Zd
„ Zd„ Zd„ Z	 dd„Z e«        e«       fd„Zd„ Zd„ Zd„ Zdedej                  dedefd„Zdedej                  dej                  dedef
d„Z dedej                  dej                  dedededefd„Z!y)é    N)Úpir)Úir_backward)Úcall_decompÚ decomp_ops_contain_unused_outputÚ
has_decomp)ÚBlockÚ	OperationÚProgram)Úcoreé   )Úregisterc                 ó¶   — t        | t        j                  «      r| fS t        | t        j                  «      rt        | «      S t        dt        | «      › d�«      S )NzType z is not supported.)Ú
isinstancer   ÚOpResultÚtypingÚSequenceÚtupleÚ	TypeErrorÚtype)Úxss    údG:\00. PROJECTS\API\Inventory\templateJSON\kerjaOCR\Lib\site-packages\paddle/decomposition/decomp.pyÚ_build_tensor_tupler       sH   € Ü�"”c—l‘lÔ#ØˆuˆÜ	�BœŸ™Ô	(Ü�R‹yÐÜ�uœT "›X˜JÐ&8Ð9Ó:Ð:ó    c                 óÊ  — t        | «      t        |«      k(  sJ ‚g }t        |«      D ]¹  \  }}t        | |   t        j                  «      r†|j                  «       t        j                  «       v r |t        |j                  «          v r|d   �/J ‚t        |«      dk(  rt        |d   t        j                  «      sJ ‚|j                  |d   «       Œ©|j                  |«       Œ» |S )Nr   r   )	ÚlenÚ	enumerater   r   r   Únamer   ÚkeysÚappend)Ú	orig_outsÚdecomp_outsÚopÚresÚidxÚvalues         r   Ú_analyse_decomp_resultsr&   (   sÄ   € Üˆy‹>œS Ó-Ò-Ð-Ð-Ø
€CÜ Ö,‰
ˆˆUÜ�i ‘n¤c§l¡lÔ3à—‘“	Ô=×BÑBÓDÑDØÔ;¸B¿G¹G»IÑFÑFà˜Q‘xÐ'Ð'Ð'ä˜5“z Q’¬:°e¸A±hÄÇÁÔ+MÐMÐMØ�J‰J�u˜Q‘xÕ à�J‰J�uÕð -ð €Jr   c                 ó4  — d}g }| j                  «       D ]¥  }|j                  «       }|r€|j                  «       rp|j                  «       }t	        |t
        «      r>|j                  «       |k(  r+|j                  «       D �cg c]  }|j                  «       ‘Œ }}|j                  |«       Œ•|j                  d«       Œ§ | j                  «       |k(  r|fS || j                  «       D �cg c]  }| j                  «       |   ‘Œ c}z   }t        |«      S c c}w c c}w )z§
    For standard api of operator, its inputs should keep consistent with organization of its inputs and attrs.

    Args:
    op (Operator): The target operator.
    úbuiltin.combineN)ÚoperandsÚsourceÚinitializedÚget_defining_opr   r	   r   r   Úget_attr_namesÚattrsr   )r"   Úcombine_op_nameÚinputsÚxÚinputÚprev_opÚitemÚapi_argumentss           r   Ú_prepare_python_api_argumentsr6   :   sù   € ð (€OØ€FØ�[‰[Ž]ˆØ—‘“
ˆÙ�U×&Ñ&Ô(Ø×+Ñ+Ó-ˆGä˜7¤IÔ.Ø—L‘L“N oÒ5à3:×3CÑ3CÔ3EÓFÑ3E¨4˜Ÿ™�Ð3E�ÐFØ�M‰M˜%Õ ð
 �M‰M˜$Õð ð" 
‡w�wƒy�OÒ#ØˆyÐà°R×5FÑ5FÔ5HÓIÑ5H°˜bŸh™h›j¨›mÐ5HÑIÑI€MÜ�ÓÐùò Gùò Js   Á?DÃ&Dc           	      ód  — d}g }| j                  «       D �]  }|j                  «       }|sŒ|j                  «       sŒ(|j                  «       }t	        |t
        «      rŒ|j                  «       |k(  ry|j                  «       D ]e  }|j                  «       j                  }d|v sŒ"t        j                  d|j                  «       j                  › d| j                  «       › d�«         y ŒÔ|j                  }d|v sŒåt        j                  d|j                  › d| j                  «       › d�«        y y )Nr(   éÿÿÿÿz;Decomp op does not support dynamic shape -1, but got shape z in inputs of op Ú Tz in op )
r)   r*   r+   r,   r   r	   r   ÚshapeÚwarningsÚwarn)r"   r/   r0   r1   r2   r3   r4   r:   s           r   Ú_check_prim_dynamicr=   [   s,  € Ø'€OØ€FØ�[‰[�]ˆØ—‘“
ˆÚ�U×&Ñ&Õ(Ø×+Ñ+Ó-ˆGä˜7¤IÔ.Ø—L‘L“N oÒ5à#×,Ñ,Ö.�DØ ŸK™K›M×/Ñ/�EØ˜U’{Ü Ÿ™ØYÐZ^×ZeÑZeÓZg×ZmÑZmÐYnÐnð  AC÷  AHñ  AHó  AJð  @Kð  KLð  Môò  $ñ /ð Ÿ™�Ø˜’;Ü—M‘MØUÐV[×VaÑVaÐUbÐbiÐjl×jqÑjqÓjsÐitÐtuÐvôñ  ñ+ r   c           	      ó,  — t        |«      t        |«      k(  s"J d| › dt        |«      › dt        |«      › �«       ‚t        ||«      D ]Ì  \  }}|�|€&| t        j                  vrt	        d| › d|› d|› �«      ‚|€Œ3|�—|�|�||j                  «       v r||||   <   |j                  }|j                  }|j                  }	|j                  }
||k(  sJ d| › d|› d	|› �«       ‚d
|
vsJ d| › d�«       ‚|	|
k(  sJ d| › d|	› d|
› �«       ‚|du |du z  rJ d«       ‚ y y)az  
    Check whether the replaced outputs are consistent with origin outputs.

    Args:
    op_name (str): The name of operator.
    orig_outs (tuple): The outputs of original operator.
    new_outs (tuple): The outputs of replaced operator.
    orig_vars (dict): Origin variables of original block.
    dst_vars (list): Corresponding replaced variables of Origin variables.
    zwhen replace origin op z[ with composite rule, num of origin outs should be equal to new outs, but len(orig_outs) = z and len(new_outs) = Nzop z2 should not contain any None value. original outs=z and its composite rule outs=z\ with composite rule, origin out dtype should be equal to new out dtype, but orig_out dtype=z and new_out dtype=r8   z1 with composite rule, composite out shape has -1.z\ with composite rule, origin out shape should be equal to new out shape, but orig_out shape=z and new_out shape=z"orig_out and new_out should match.)r   Úzipr   Úops_contain_noneÚ
ValueErrorr   Údtyper:   )Úop_namer    Únew_outsÚ	orig_varsÚdst_varsÚorig_outÚnew_outÚ
orig_dtypeÚ	new_dtypeÚ
orig_shapeÚ	new_shapes              r   Ú_check_op_resultsrM   v   sÊ  € ô ˆy‹>œS ›]Ò*ð Ø
! ' ð + Ü # I£Ð/Ð/DÄSÈÃ]ÀOð	UóÐ*ô
 !ØØöÑˆ�'ð Ð  Øœ4×0Ñ0Ñ0äØ�g�YÐPÐQZÐP[Ð[xð  zBð  yCð  Dóð ð ÐàØÐ ØÐ$¨Ð)=Ø˜yŸ~™~Ó/Ñ/Ø4;�H˜Y xÑ0Ñ1Ø!Ÿ™ˆJØŸ™ˆIØ!Ÿ™ˆJØŸ™ˆIØ Ò*ð Ø)¨'¨ð 3&Ø&0 \Ð1DÀYÀKðQóÐ*ð
 ˜)Ñ#ðdà(¨¨	Ð1bÐcódØ#à Ò*ð Ø)¨'¨ð 3&Ø&0 \Ð1DÀYÀKðQóÐ*ð ! DÐ(Ø˜4�òð 4à3ó4ð ñ 	ñGr   c                 óØ  ‡‡— t        j                  «       s|S t        | t        «      st	        dt        | «      › d�«      ‚| j                  «       }t        ‰t        t        f«      st	        dt        ‰«      › d�«      ‚t        ‰t        t        f«      st	        dt        ‰«      › d�«      ‚t         j                  d   ‰z  Št        j                  d«       t        ‰«      dkD  rt        ‰«      dkD  rˆˆfd„}nGt        ‰«      dkD  rt        ‰«      dk(  rˆfd	„}n%t        ‰«      dk(  rt        ‰«      dkD  rˆfd
„}nd„ }dgt        |«      z  }i }t        |«      D ]?  \  }}	t        |	t        j                  «      st	        dt        |	«      › d|› d�«      ‚|||	<   ŒA t        j                   j!                  | «      5  t#        ||||«       ddd«       t        |«      D ]E  \  }}	t        |	t        j                  «      rŒ!|	€	||   ||<   Œ,t	        dt        |	«      › d|› d�«      ‚ t        j                  dj%                  t         j                  d   «      «       |S # 1 sw Y   Œ“xY w)a¢  
    Search nonbasic ops which have be registered composite rules and replace them with primitive ops.
    The operators in blacklist will be excluded from program when decomposed into primitives, and only the
    operators in whitelist will be decomposed. The priority of blacklist is higher than whitelist, it means
    an operator both in blacklist and whitelist will not be decomposed.

    The finally set that will be decomposed is:
        (block.ops & ops have decomposite rule & whitelist) - blacklist

    Note:
        All variables must be contained inside the given program.

    Args:
        program (Program): The program to be processed.
        src_vars (list[OpResult]): In program, once some operator is decomposed, its vars will be replaced by new ones. This argument means some vars will be used later and corresponding vars will be returned for later usage.
        blacklist (frozenset): The Operators that will be exclude when decomposed into primitives.
        whitelist (frozenset): Only the operators in whitelist will be decomposed into primitives.

    Returns:
        dst_vars (list): A list contains all vars which replace origin ones in src_vars.
    z"Expect type Program, but got type Ú.z6Expected type of blacklisst is set|frozenset, but got z6Expected type of whiltelist is set|frozenset, but got Úforward_blacklistz(Decompose composite forward ops begin...r   c                 óP   •— | j                  «       ‰v xr | j                  «       ‰vS ©N©r   )r1   Ú	blacklistÚ	whitelists    €€r   Ú<lambda>zdecompose.<locals>.<lambda>Þ   s#   ø€ �a—f‘f“h )Ð+ÒI°·±³À	Ð0IÐIr   c                 ó(   •— | j                  «       ‰vS rR   rS   )r1   rT   s    €r   rV   zdecompose.<locals>.<lambda>á   s   ø€ ˜aŸf™f›h¨iÑ7r   c                 ó(   •— | j                  «       ‰v S rR   rS   )r1   rU   s    €r   rV   zdecompose.<locals>.<lambda>ã   s   ø€ ˜aŸf™f›h¨)Ñ3r   c                  ó   — y)NT© )r1   s    r   rV   zdecompose.<locals>.<lambda>å   s   € ˜dr   NzLEach var in dst_vars should map corresponding var in src_vars, but got type z in z*Decompose composite forward ops finish: {}Úcomposite_ops_record)r   Ú_is_fwd_prim_enabledr   r
   r   r   Úglobal_blockÚsetÚ	frozensetÚprim_configÚloggingÚdebugr   r   r   r   Úprogram_guardÚ_decompose_subgraphÚformat)
ÚprogramÚsrc_varsrT   rU   ÚblockÚ	op_filterrF   Údst_vars_dctr$   r4   s
     ``      r   Ú	decomposerk   ®   s]  ù€ ô6 ×$Ñ$Ô&ØˆÜ�gœwÔ'ÜÐ<¼TÀ'»]¸OÈ1ÐMÓNÐNØ× Ñ Ó"€Eä�i¤#¤yÐ!1Ô2ÜØDÄTÈ)Ã_ÐDUÐUVÐWó
ð 	
ô �i¤#¤yÐ!1Ô2ÜØDÄTÈ)Ã_ÐDUÐUVÐWó
ð 	
ô × Ñ Ð!4Ñ5¸	ÑA€Iä‡M�MÐ<Ô=ä
ˆ9ƒ~˜Òœc )›n¨qÒ0äIñ 	ô 
ˆY‹˜!Ò	¤ I£°!Ò 3Û7‰	Ü	ˆY‹˜1Ò	¤ Y£°!Ò!3Û3‰	á"ˆ	Øˆvœ˜H›Ñ%€HØ€LÜ˜xÖ(‰	ˆˆTÜ˜$¤§¡Ô-ÜØ^Ô_cÐdhÓ_iÐ^jÐjnÐowÐnxÐxyÐzóð ð !ˆ�TÒð )ô 
�‰×	Ñ	 Õ	(ÜØØØØô		
÷ 
)ô ˜xÖ(‰	ˆˆTÜ˜$¤§¡Õ-Øˆ|Ø (¨¡�˜’äØbÔcgÐhlÓcmÐbnÐnrÐs{Ðr|Ð|}Ð~óð ð )ô ‡M�MØ4×;Ñ;Ü×ÑÐ3Ñ4ó	
ôð
 €O÷) 
)Ð	(ús   Æ?I É I)c                 óˆ  — t        | t        «      �rk| j                  }d}t        |«      D �]M  \  }}|j	                  «       }t        j                  |«      }	t        |«      }
|	xs |
xr  ||«      }|r!t        j                  «       rt        |«      rd}|j	                  «       dk(  r|}|sŒ‚t        j                  d   j                  |«       |�/||dz
     j	                  «       dk(  rt        j                  |«       nt        j                  |«       t        |«      }|j!                  «       }|
rt#        |«      }t%        |||«      }nt'         |	|Ž «      }t)        |||||«       |j	                  «       t+        j,                  «       v rKt/        t1        |«      «      D ]3  }|t*        |j	                  «          vsŒ||   j3                  ||   «       Œ5 nM|j	                  «       t+        j,                  «       v r|d   j3                  |d   «       n|j3                  |«       | j5                  |«       |€�Œd}|j!                  «       D ]  }|j7                  «       sŒd} n |r| j5                  |«       d}�ŒP yt        | t8        j:                  «      r| D ]  }t=        ||||«       Œ yt?        dtA        | «      › �«      ‚)	aš  
    The operators in block wich satisfy the filter conditon will be decomposed into primitives.

    Args:
        block (Block|Sequence[Block]): The blocks of program to be processed.
        op_filter (function): The filter to specify which ops to be processed.
        orig_vars (dict): Origin variables of original block.
        dst_vars (list): Corresponding replaced variables of Origin variables.
    NFr(   r[   r   r   Tz5Expect type Block or Sequence of Block, but got type )!r   r   Úopsr   r   r   Úget_decomp_ruler   r   Ú_enable_prim_dynamic_shaper=   r`   Úaddr   Úset_insertion_pointr6   Úresultsr   r&   r   rM   r   r   Úranger   Úreplace_all_uses_withÚ	remove_opÚhas_one_user   r   rd   r   r   )rh   rE   rF   ri   Úops_listÚtemp_opr$   r"   rC   Ú
decom_ruleÚhas_sink_decomp_ruleÚlowerÚ
input_argsr    r!   rD   ru   r4   s                     r   rd   rd     sƒ  € ô �%œÕØ—9‘9ˆØˆÜ  ×*‰GˆC�Ø—g‘g“iˆGÜ!×1Ñ1°'Ó:ˆJÜ#-¨b£>Ð ØÒ7Ð#7ÒJ¹YÀr»]ˆEñ Ü×3Ñ3Ô5Ü'¨Ô+à�à�w‰w‹yÐ-Ò-Ø�âÜ× Ñ Ð!7Ñ8×<Ñ<¸WÔEàÐ'Ø   q¡Ñ)×.Ñ.Ó0Ð4EÒEä×+Ñ+¨GÕ4ä×+Ñ+¨BÔ/Ü:¸2Ó>�
ØŸJ™J›L�	Ù'Ü"-¨b£/�KÜ6Ø! ;°ó ‘Hô  3±:¸zÐ3JÓK�Hô "Ø˜Y¨°)¸Xôð —7‘7“9Ô @× EÑ EÓ GÑGÜ$¤S¨£^Ö4˜àÜ#CÀBÇGÁGÃIÑ#NòOð & c™N×@Ñ@ÀÈ#ÁÕOñ  5ð —w‘w“yÔ$D×$IÑ$IÓ$KÑKØ! !™×:Ñ:¸8ÀA¹;ÕGà×0Ñ0°Ô:Ø—‘ Ô#àÒ&Ø $�IØ '§¡Ö 1˜Ø×+Ñ+Õ-Ø(-˜IÙ!ð !2ñ !ØŸ™¨Ô0Ø"’Gð{ +ð| 	ä	�Eœ6Ÿ?™?Ô	+ÛˆDÜ  i°¸9ÕEð àÜ
Ø
?ÄÀUÃ¸}ÐMóð r   c                 óò  — t        | t        «      st        dt        | «      › �«      ‚t        |t        «      st        d«      ‚g }g }i }| j
                  D ]@  }g }|j                  «       D ]$  }|j                  «       sŒ|j                  |«       Œ& |||<   ŒB |D ]Æ  }|j                  «       D ]±  }|||   v sŒ||vrR|j                  |«       |j                  |j                  «       j                  |«      |j                  |«      gg«       Œa||j                  |«         j                  |j                  «       j                  |«      |j                  |«      g«       Œ³ ŒÈ t        |«      t        |«      fS )aP  
    This API checks which op contributes to the outputs of the entire computation graph,
    as well as determining the corresponding output index.

    Args:
        block (Block): the block of program to be processed.
        global_outputs (tuple(Value)): the outputs of the entire computation graph.

    Returns:
        related_ops (tuple(pir.Operation)): a tuple of op that contributes to the outputs of the entire graph.
        related_ops_output_indexes (tuple(tuple())) : a tuple records the mapping of tuple(the output index of the op,  the output index of the entire graph)
    z$block should be Block, but got type z)The type of global_outputs should be list)r   r   r   r   Úlistrm   rr   r+   r   r   Úindexr   )	rh   Úglobal_outputsÚrelated_opsÚrelated_ops_output_indexesÚop_to_op_valid_resultr"   Úop_valid_resultr1   Úglobal_outputs	            r   Úget_leaf_opsr†   \  sq  € ô �eœUÔ#ÜÐ>¼tÀE»{¸mÐLÓMÐMÜ�n¤dÔ+ÜÐCÓDÐDà€KØ!#ÐàÐØ�iŒiˆØˆØ—‘–ˆAØ�}‰}�Ø×&Ñ& qÕ)ð ð %4Ð˜bÒ!ð ó (ˆØ'×,Ñ,Ö.ˆBØÐ 5°bÑ 9Ò9Ø˜[Ñ(Ø×&Ñ& rÔ*Ø.×5Ñ5ð !#§
¡
£× 2Ñ 2°=Ó AØ .× 4Ñ 4°]Ó Cððõð /¨{×/@Ñ/@ÀÓ/DÑE×LÑLàŸJ™J›L×.Ñ.¨}Ó=Ø*×0Ñ0°Ó?ðõñ /ð (ô* �ÓœuÐ%?Ó@Ð@Ð@r   c                 ó4   — ||   D ]  }||d      | |d   <   Œ y)z²
    This API replace the outputs of the entire computation graph with the new outputs of the op,
    when the op contributes to the outputs of the entire computation graph.
    r   r   NrZ   )r€   Ú
op_outputsÚop_indexr‚   r   s        r   Úreplace_graph_outputsrŠ   ‘  s*   € ð ,¨HÔ5ˆØ#-¨e°A©hÑ#7ˆ�u˜Q‘xÒ ñ 6r   rh   Úfwd_opÚgrad_var_to_var_mapÚreturnc                 óÚ  — t        j                  «       st        d«      ‚t        j                   j	                  | j
                  «      5  |j                  «       }|j                  «       }t        j                  |«      }t        |«      }|xs |}|r¹t        |«      }t        j                  |«       |rt        |«      }	t        ||	|«      }
nt         ||Ž «      }
t!        |||
«       |j#                  «       D ]!  \  }}||v sŒ|
|j%                  |«         ||<   Œ# |j'                  |
«       | j)                  |«       |
dfcddd«       S t+        |«      dfcddd«       S # 1 sw Y   yxY w)a  
    Decompose the fwd_op into a list of primitive ops.

    Args:
        block (Block): the block to which the fwd_op belongs.
        fwd_op (pir.Operation): the forward op to be decomposed.
        grad_var_to_var_map (dict): a dict obtained from distributed processing,
            which maps the backward grad variable to its corresponding forward variable.
    Returns:
        new_outputs (tuple(Value)): the new outputs after decomposing.
        has_decomposed: whether the forward op has been successfully decomposed.
    zRTo decompose forward op, please set `core._set_prim_forward_enabled(True)` firstlyTNF)r   r\   ÚRuntimeErrorr   rc   rf   r   rr   r   rn   r   r6   rq   r   r&   r   rM   Úitemsr   rt   ru   r   )rh   r‹   rŒ   rC   r    ry   rz   r{   r|   r!   rD   Úgrad_varÚvars                r   Údecompose_fwd_opr“   Ÿ  sN  € ô  ×$Ñ$Ô&ÜØ`ó
ð 	
ô 
�‰×	Ñ	 §¡Õ	.Ø—+‘+“-ˆØ—N‘NÓ$ˆ	Ü×-Ñ-¨gÓ6ˆ
Ü)¨&Ó1ÐØÒ2Ð2ˆáÜ6°vÓ>ˆJÜ×#Ñ# FÔ+Ù#Ü)¨&Ó1�Ü2Ø˜{¨Fó‘ô /©z¸:Ð/FÓG�ä˜g y°(Ô;ð "5×!:Ñ!:Ö!<‘�˜#Ø˜)Ò#Ø4<Ø!Ÿ™¨Ó,ñ5Ð'¨Ò1ð "=ð ×(Ñ(¨Ô2Ø�O‰O˜FÔ#Ø˜T�>÷; 
/Ñ	.ô> ˜Ó# UÐ*÷? 
/×	.Ò	.ús   Á	B9E!Ä>E!ÅE!Å!E*Úbwd_opc                 ó
  — t        j                  «       st        d«      ‚|j                  «       D �cg c]  }|j	                  «       ‘Œ }}|j                  «       }|j                  «       D �cg c]  }|j	                  «       ‘Œ }}|j                  «       }g }	g }
|D ]  }||v rŒ||v rŒ|
j                  |g«       Œ  |D �cg c]  }|g‘Œ }}t        d|j                  «       «      D �cg c]  }|j                  |«      g‘Œ }}g }|D ]7  }|j                  «       r|j                  dg«       Œ&|j                  dg«       Œ9 | j                  j                  |«      }t        | j                  «      }t        j                  ||||
|«      }t        | j                  «      }||z
  }|dk(  rM| j                  d   j                  «       |j                  «       k(  r| j!                  | j                  d   «       y|D ]R  }|d   �(|d   j                  «       r|	j                  |d   «       Œ0|	j                  t#        j$                  «       «       ŒT t'        |«      D ]/  \  }}||j)                  «       v sŒ|j+                  |«      ||	|   <   Œ1 |}t        ||«      D ]&  }| j-                  | j                  |   |«       |dz  }Œ( |j/                  |	«       | j!                  |«       t1        |	«      dfS c c}w c c}w c c}w c c}w )ac  
    Decompose the bwd_op into a list of primitive ops.
    If fwd_op has composite vjp rules (including custom vjp), call call_vjp() to get a list of primitive operators in backward graph, then replace bwd_op.

    Args:
        block (Block): the block to which the bwd_op belongs.
        fwd_op (pir.Operation): the forward op.
        bwd_op (pir.Operation): the backward op to be decomposed.
        grad_var_to_var_map (dict): a dict obtained from distributed processing,
            which maps the backward grad variable to its corresponding forward variable.
    Return:
        new_input_grads (tuple(Value)): new results of backward op after decomposing.
        has_decomposed: whether the backward op has been successfully decomposed. If a fwd op does not have composite vjp rules and can not be decomposed directly, this function will return False.
    úTTo decompose backward op, please set `core._set_prim_backward_enabled(True)` firstlyr   FTr   r8   )NF)r   Ú_is_bwd_prim_enabledr�   r)   r*   rr   r   rs   Únum_operandsÚoperand_sourcer+   rm   r   r   Úcall_vjpr   ru   r   Úfake_op_resultr   r   ÚpopÚmove_oprt   r   )rh   r‹   r”   rŒ   r1   Ú
fwd_inputsÚfwd_outputsÚ
bwd_inputsÚgrad_inputsr#   Úgrad_outputsÚ	bwd_inputÚ
fwd_outputÚfwd_outputs_ÚiÚfwd_inputs_Ústop_gradientsÚ
grad_inputÚ
bwd_op_idxÚbefore_num_opsÚnew_grad_inputsÚafter_num_opsÚnum_appended_opsr$   Ú
insert_idxs                            r   Údecompose_bwd_op_directlyr°   Ö  sÜ  € ô* ×$Ñ$Ô&ÜØbó
ð 	
ð
 '-§o¡oÔ&7Ó8Ñ&7 �!—(‘(•*Ð&7€JÐ8Ø—.‘.Ó"€KØ&,§o¡oÔ&7Ó8Ñ&7 �!—(‘(•*Ð&7€JÐ8Ø—.‘.Ó"€KØ
€Cð €LÛˆ	Ø˜ZÒ'¨9¸Ò+CØ×Ñ  Õ,ð  ñ 4?Ó?±; Z�Z’L°;€LÐ?ä,1°!°V×5HÑ5HÓ5JÔ,KóÙ,K qˆ×	Ñ	˜qÓ	!Ò"Ð,Kð ð ð €NÛ!ˆ
Ø×!Ñ!Ô#Ø×!Ñ! 5 'Õ*à×!Ñ! 4 &Õ)ð	 "ð —‘—‘ Ó(€JÜ˜Ÿ™“^€Nä—m‘mØ�˜\¨<¸ó€Oô ˜Ÿ	™	“N€MØ$ ~Ñ5Ðð ˜1Ò §¡¨2¡×!3Ñ!3Ó!5¸¿¹»Ò!FØ�‰˜Ÿ	™	 "™Ô&Øó *ˆJØ˜!‰}Ð(¨Z¸©]×-FÑ-FÔ-HØ—
‘
˜: a™=Õ)à—
‘
œ3×-Ñ-Ó/Õ0ð	 *ô  )¨Ö5‰OˆC�ØÐ0×5Ñ5Ó7Ò7Ø0C×0GÑ0GØó1Ð# C¨¡HÒ-ð  6ð  ˆ
Ü�~ }Ö5ˆAØ�M‰M˜%Ÿ)™) A™,¨
Ô3Ø˜!‰O‰Jð 6ð
 	×$Ñ$ SÔ)Ø�‰˜Ôä�S‹z˜4ÐÐùò} 9ùâ8ùò @ùòs   ²K1Á-K6Ã
K;Ã)L rž   Úfwd_outputs_after_decomposec                 ó  ‡‡‡— t        j                  «       st        d«      ‚‰€t        d«      ‚|j                  «       D �cg c]  }|j	                  «       ‘Œ }}|j                  «       }g }	t        ˆˆfd„|D «       «      }
t        ˆfd„|
D «       «      }t        ˆfd„|D «       «      }| j                  j                  |«      }t        | j                  «      }t        j                  |||
«      }t        | j                  «      }d}t        |«      D ]R  \  }}|j                  «       r|	j                  ||   «       |dz  }Œ0|	j                  t        j                   «       «       ŒT t        |«      D ]/  \  }}|‰j#                  «       v sŒ‰j%                  |«      ‰|	|   <   Œ1 |}t'        ||«      D ]&  }| j)                  | j                  |   |«       |dz  }Œ( |j+                  |	«       | j-                  |«       t        |	«      S c c}w )a”  
    Decompose the bwd_op into a list of primitive ops.
    If fwd_op has no composite vjp rules, and fwd_op has been decomposed to a list of primitive operators in forward graph previously,
    call grad() for the decomposed forward subgraph to get a list of primitive operators in backward graph, then replace bwd_op.

    Args:
        block (Block): the block to which the bwd_op belongs.
        fwd_op (pir.Operation): the forward op.
        bwd_op (pir.Operation): the backward op to be decomposed.
        grad_var_to_var_map (dict): a dict obtained from distributed processing,
            which maps the backward grad variable to its corresponding forward variable.
        fwd_inputs: (tuple(Value)): the original input of the forward op,
        fwd_outputs_after_decompose (tuple(Value)): the output of the decomposed forward op, if forward op has no vjp rules, forward op shoule be decomposed firstly,
            fwd_outputs_after_decompose means the new output of the decomposed forward op. If forward op has vjp rules, fwd_outputs_after_decompose is None.
    Return:
        new_input_grads (tuple(Value)): results of backward op after decomposing.
    r–   z=To decompose backward op, please decompose forward op firstlyc              3   ó2   •K  — | ]  }|‰v s|‰v s|–— Œ y ­wrR   rZ   )Ú.0r£   rž   r±   s     €€r   Ú	<genexpr>z0decompose_bwd_op_after_fwd_op.<locals>.<genexpr>[  s)   øè ø€ ð á#ˆIà˜Ñ# yÐ4OÑ'Oô 	Ù#ùs   ƒc              3   ó(   •K  — | ]	  }‰|   –— Œ y ­wrR   rZ   )r´   Úgrad_outputrŒ   s     €r   rµ   z0decompose_bwd_op_after_fwd_op.<locals>.<genexpr>b  s   øè ø€ ð Ù<H¨[Ð˜KÕ(¹Lùs   ƒc              3   óH   •K  — | ]  }|j                  «       r‰|   –— Œ y ­wrR   )r+   )r´   r©   rŒ   s     €r   rµ   z0decompose_bwd_op_after_fwd_op.<locals>.<genexpr>e  s*   øè ø€ ð á%ˆJØ×!Ñ!Ô#ð 	˜JÕ'Ù%ùs   ƒ"r   r   )r   r—   r�   r)   r*   rr   r   rm   r   r   r   Úgradr   r+   r   r   r›   r   rœ   rs   r�   rt   ru   )rh   r‹   r”   rŒ   rž   r±   r1   r    r¡   r#   r¢   r¥   r§   rª   r«   r¬   r­   Úinput_grads_idxr$   r©   r¯   r¦   s      ```                r   Údecompose_bwd_op_after_fwd_opr»   2  sö  ú€ ô4 ×$Ñ$Ô&ÜØbó
ð 	
ð #Ð*ÜØKó
ð 	
ð
 '-§o¡oÔ&7Ó8Ñ&7 �!—(‘(•*Ð&7€JÐ8Ø—.‘.Ó"€KØ
€Cô ô á#óó €Lô ó Ù<Hóó €Lô ó á%óó €Kð —‘—‘ Ó(€JÜ˜Ÿ™“^€Nä!×&Ñ& |°[À,ÓO€OÜ˜Ÿ	™	“N€Mð €OÜ$ [Ö1‰ˆˆZØ×!Ñ!Ô#Ø�J‰J� Ñ7Ô8Ø˜qÑ ‰Oà�J‰J”s×)Ñ)Ó+Õ,ð 2ô % [Ö1‰ˆˆZØÐ,×1Ñ1Ó3Ò3Ø,?×,CÑ,CÀJÓ,OÐ  C¡Ò)ð 2ð
 €JÜ�> =Ö1ˆØ�‰�e—i‘i ‘l JÔ/Ø�a‰‰
ð 2ð
 × Ñ  Ô%Ø	‡O�O�FÔä�‹:Ðùòi 9s   ÁH	)NN)"ra   r   r;   Úpaddler   Úpaddle.autogradr   Úpaddle.base.corer   r   r   Úpaddle.base.libpaddle.pirr   r	   r
   Úpaddle.frameworkr   Ú r   r   r&   r6   r=   rM   r_   rk   rd   r†   rŠ   Údictr   r“   r°   r»   rZ   r   r   Ú<module>rÃ      s9  ðó Û Û å Ý '÷ñ ÷
 @Ñ ?Ý !å ò;òò$ òB ð8 <@ó5ñv ‹kÙ‹kó	TònTòn2Aòj8ð4+Øð4+ØŸ-™-ð4+Ø>Bð4+à
ó4+ðnY ØðY à�M‰MðY ð �M‰MðY ð ð	Y ð
 óY ðxXØðXà�M‰MðXð �M‰MðXð ð	Xð
 ðXð "'ðXð ôXr   