a
    d3                     @   s   d dl Z d dlZd dlZd dlmZmZmZmZ d dlZ	ddl
mZmZ ddlmZ ddlmZ ddlmZmZ ddlmZ dd	lmZ ejG d
d dZejG dd dZG dd deZejG dd deZG dd dZdS )    N)AnyDictListOptional   )utils	variables)create_instruction)	PyCodegen)LocalSourceSource)
object_new)VariableTrackerc                   @   s6   e Zd ZU dZeed< dZeed< dd Zdd Z	d	S )
MutableSideEffectsz
    VariableTracker.mutable_local marker to indicate a list passed as
    an input that if we mutate we need to re-apply those mutations after
    the graph runs.
    sourceFis_modifiedc                 C   s   t | S Nidself r   c/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/torch/_dynamo/side_effects.py__hash__   s    zMutableSideEffects.__hash__c                 C   s   | |u S r   r   r   otherr   r   r   __eq__   s    zMutableSideEffects.__eq__N)
__name__
__module____qualname____doc__r   __annotations__r   boolr   r   r   r   r   r   r      s
   
r   c                   @   s   e Zd ZU dZeed< dS )AttributeMutationzM
    VariableTracker.mutable_local marker to track changes to attributes
    r   N)r   r   r   r    r   r!   r   r   r   r   r#   "   s   
r#   c                   @   s   e Zd Zdd Zdd ZdS )AttributeMutationExistingc                 C   s   t | S r   r   r   r   r   r   r   ,   s    z"AttributeMutationExisting.__hash__c                 C   s   | |u S r   r   r   r   r   r   r   /   s    z AttributeMutationExisting.__eq__N)r   r   r   r   r   r   r   r   r   r$   +   s   r$   c                   @   s&   e Zd ZU eed< dd Zdd ZdS )AttributeMutationNew
cls_sourcec                 C   s   t | S r   r   r   r   r   r   r   7   s    zAttributeMutationNew.__hash__c                 C   s   | |u S r   r   r   r   r   r   r   :   s    zAttributeMutationNew.__eq__N)r   r   r   r   r!   r   r   r   r   r   r   r%   3   s   
r%   c                       s  e Zd ZU dZeeef ed< eeee	ef f ed< e
e ed< dH fdd	Zeedd	d
Zd ee	 dddZdd Zddd fddZdd Zdd Zee	edddZdd Zdd Zdd Zee	d d!d"Zee	ed#d$d%Zed&d' Zd(d) Zd*d+ Ze fe!eed,d-d.Z"e"Z#e"Z$e!eed,d/d0Z%e!eed1d2d3Z&d4d5 Z'e!ed6d7d8Z(e!ed6d9d:Z)d;d< Z*d=d> Z+d?d@ Z,e-dAdBdCZ.e-dAdDdEZ/dFdG Z0  Z1S )ISideEffectszp
    Track side effects (list mutation, setattr, etc) that need to be
    applied after an FX graph is run.
    id_to_variablestore_attr_mutations	keepaliveNc                    s4   t    |pt | _|p"t | _|p,g | _d S r   )super__init__collectionsOrderedDictr(   r)   r*   )r   r(   r)   r*   	__class__r   r   r,   H   s    
zSideEffects.__init__)r   returnc                 C   s&   t |tsJ | j|jko$| j|jkS r   )
isinstancer'   r(   r)   r   r   r   r   r   N   s    
zSideEffects.__eq__c                 C   s   | j |j kr<| j  }|j  }||kr8d| d| S dS | j|jkrx| j }|j }||krtd| d| S dS d S d S )Nzid_to_variable keys: z != zid_to_variable: unknown diffzstore_attr_mutations keys: z"store_attr_mutations: unknown diff)r(   keysr)   )r   r   Zsk_itvZok_itvZsk_samZok_samr   r   r   diffV   s    



zSideEffects.diffc                 C   s4   | j t| jtdd | j D t| jdS )zCreate a shallow copyc                 s   s    | ]\}}|t |fV  qd S r   )r-   r.   .0kvr   r   r   	<genexpr>l   s   z$SideEffects.clone.<locals>.<genexpr>)r(   r)   r*   )r0   r-   r.   r(   r)   itemslistr*   r   r   r   r   cloneh   s    

zSideEffects.clonec                 C   s   dS )NFr   )_r   r   r   <lambda>s       zSideEffects.<lambda>c                    sZ    d u rt   t fdd| j D | _t fdd| j D | _d S )Nc                 3   s&   | ]\}}|t | fV  qd S r   r   applyr5   cachefnskip_fnr   r   r9   w   s   z$SideEffects.apply.<locals>.<genexpr>c                 3   s&   | ]\}}|t | fV  qd S r   r@   r5   rB   r   r   r9   {   s   )dictr-   r.   r(   r:   r)   )r   rD   rC   rE   r   rB   r   rA   s   s    
zSideEffects.applyc                 C   s   t || jv S r   )r   r(   r   itemr   r   r   __contains__   s    zSideEffects.__contains__c                 C   s   | j t| S r   )r(   r   rG   r   r   r   __getitem__   s    zSideEffects.__getitem__)rH   namevaluec                 C   s>   |  |sJ |j| jvr*t | j|j< || j|j |< d S r   )is_attribute_mutationmutable_localr)   r-   r.   )r   rH   rK   rL   r   r   r   
store_attr   s    zSideEffects.store_attrc                 C   s   |  |sJ | j|j | S r   )rM   r)   rN   )r   rH   rK   r   r   r   	load_attr   s    zSideEffects.load_attrc                 C   s2   t |tjsJ t |tjs J | |d| d S Ncell_contents)r2   r   NewCellVariabler   rO   )r   cellvarrL   r   r   r   
store_cell   s    zSideEffects.store_cellc                 C   s   t |tjsJ | |dS rQ   )r2   r   rS   rP   )r   rT   r   r   r   	load_cell   s    zSideEffects.load_cell)gvarrK   c                 C   s   t |tjsJ | ||S r   )r2   r   r   rP   )r   rW   rK   r   r   r   load_global   s    zSideEffects.load_global)rW   rK   rL   c                 C   s2   t |tjsJ t |tjs J | ||| d S r   )r2   r   r   rO   )r   rW   rK   rL   r   r   r   store_global   s    zSideEffects.store_globalc                 C   s   t | dd tjtjjjfv S )N__setattr__)inspectgetattr_staticobjectrZ   torchnnModule)clsr   r   r   "cls_supports_mutation_side_effects   s    z.SideEffects.cls_supports_mutation_side_effectsc                 C   s   t |jtS r   )r2   rN   r#   rG   r   r   r   rM      s    z!SideEffects.is_attribute_mutationc                 C   s.   t |jtrdS | |r&|j| jv S |jjS NT)r2   rN   r%   rM   r)   r   rG   r   r   r   r      s
    
zSideEffects.is_modified)r   rH   variablec                 C   s0   |j |||d}|| jt|< | j| |S )z*Start tracking a new variable for mutation)rN   r   )r<   r(   r   r*   append)r   r   rH   rd   mutable_clsr   r   r   
_track_obj   s    zSideEffects._track_objc                 C   s   | j |||tdS )N)rf   )rg   r$   r   r   rH   rd   r   r   r   track_object_existing   s    z!SideEffects.track_object_existing)r&   user_clsvariable_clsc                 C   s@   t |}||fdtd |i|}|| jt|< | j| |S )NrN   )r   r%   r(   r   r*   re   )r   r&   rj   rk   optionsobjrd   r   r   r   track_object_new   s    zSideEffects.track_object_newc                 C   s6   t  }tjtd d d}|| jt|< | j| |S NrN   )r]   r   rS   r%   r(   r   r*   re   )r   rm   rd   r   r   r   track_cell_new   s    zSideEffects.track_cell_new)r   rH   c                 C   s.   t jt|d}|| jt|< | j| |S ro   )r   rS   r$   r(   r   r*   re   rh   r   r   r   track_cell_existing   s    zSideEffects.track_cell_existingc                 C   s.   t jt|d}|| jt|< | j| |S ro   )r   NewGlobalVariabler$   r(   r   r*   re   rh   r   r   r   track_global_existing   s    z!SideEffects.track_global_existingc                    s   t  d tdfdd}td fdd t||j|jf | j D ]}t|jt	sPt|| qP| j
 D ]\}t|| qxt fdd| j D | _t fdd| j
 D | _
d S )	Nvarc                    s&   t | jtr"| jur" | j | S r   )r2   rN   r%   addru   )live_new_objectsskip_objr   r   visit   s    
z0SideEffects.prune_dead_object_new.<locals>.visitc                    s*   t | tr| v S t | tr& | jS dS rc   )r2   r%   r   rN   ru   )is_liverx   r   r   r{     s
    


z2SideEffects.prune_dead_object_new.<locals>.is_livec                 3   s"   | ]\}} |r||fV  qd S r   r   r5   r{   r   r   r9     s   z4SideEffects.prune_dead_object_new.<locals>.<genexpr>c                 3   s"   | ]\}} |r||fV  qd S r   r   r5   r|   r   r   r9     s   )setr   rA   stackZsymbolic_localsr(   valuesr2   rN   r%   r)   r:   r-   r.   )r   txrz   rv   Zsetattrsr   )r{   rx   ry   r   prune_dead_object_new   s     
z!SideEffects.prune_dead_object_newc                 C   s   |j t|jjddS )NTrp   )r<   r   rN   r   )r   ZoldvarZnewvarr   r   r   mutation  s    zSideEffects.mutationc                    s    fdd j  D S )Nc                    s   g | ]}  |r|qS r   )r   )r6   rv   r   r   r   
<listcomp>!  r?   z2SideEffects._get_modified_vars.<locals>.<listcomp>)r(   r   r   r   r   r   _get_modified_vars   s    zSideEffects._get_modified_vars)cgc                 C   s  |   D ]}t|jttfrrt|tjrr|tj	d |
tddg || t|jtrt|j| |j_qt|jtr|tj	d ||jj |
tddg || t|j| |j_q||jv r|j|d u sJ ||jj || qd S )NZ	make_cellCALL_FUNCTIONr   r   r   )r   r2   rN   r$   r%   r   rS   Zload_import_fromr   r   extend_outputr	   Z	add_cacher   Ztempvarsr   r&   get)r   r   rv   r   r   r   codegen_save_tempvars#  s*    




z!SideEffects.codegen_save_tempvarsc                 C   s  g }|   D ]}t|tjrj||dd ||jj ||d |d tddg |	tdg qt|tj
r|jjd |jjd ||jj |tddg ||dd ||jj |tddg |	td	d
tdtd	dtdg q| |r| j|ji  D ]v\}}t|tjrd|jj| || |	td|g n4|jj| || ||jj |	td|g q$qtt|qt|D ]}|| qd S )NF)Zallow_cacheBUILD_SLICE   STORE_SUBSCRclearupdateLOAD_METHODCALL_METHODr   POP_TOPr   STORE_GLOBAL
STORE_ATTR)r   r2   r   ZListVariablerN   r   r   Zcreate_load_constr	   re   ZConstDictVariabler   outputZupdate_co_namesrM   r)   r   r:   rs   AssertionErrortypereversed)r   r   suffixesrv   rK   rL   suffixr   r   r   codegen_update_mutated9  sT    z"SideEffects.codegen_update_mutatedc                 C   s   t t| j| j  S r   )anymapr   r(   r   r   r   r   r   is_emptyo  s    zSideEffects.is_empty)NNN)2r   r   r   r    r   intr   r!   r#   strr   r   r,   r]   r"   r   r   r4   r<   rA   rI   rJ   rO   rP   rU   rV   rX   rY   staticmethodrb   rM   r   r   r   rg   Z
track_listZ
track_dictri   rn   rq   rr   rt   r   r   r   r
   r   r   r   __classcell__r   r   r/   r   r'   >   s\   

"6r'   )r-   Zdataclassesr[   typingr   r   r   r   Ztorch.nnr^    r   r   Zbytecode_transformationr	   Zcodegenr
   r   r   r   r   Zvariables.baser   Z	dataclassr   r#   r$   r%   r'   r   r   r   r   <module>   s$   
