a
    d7                     @   sF  d dl Z d dlZ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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 d dlmZ dd	lmZmZmZ dd
lmZ ddlmZm Z m!Z! d dl"m#  m$Z% e&e'Z(dd Z)edd Z*eej+e
dddZ,d0ddZ-dd Z.eej+e
dddZ/G dd dej0Z1eej+e
dddZ2edd Z3d1d d!Z4ej5j6Z6e6j7e6j8e6j9e6j:e6j;e6j<e6j=e6j>e6j?e6j@e6jAe6jBe6jCe6jDjEe6jDjFe6jGe6jHe6jIe6jJe6jKe6jLe6jMhZNeeNZNed"d# ZOd2ee
ejPf eeeQ  d$d%d&ZRd'd( ZSd aTd)d* ZUd+d, ZVd3d.d/ZWdS )4    N)contextmanager)partial)CallableOptionalTupleUnion)SymInt)get_decompositions)bind_symbols   )aot_function
aot_modulemake_boxed_compiler)strip_overloads)default_partition
draw_graph#min_cut_rematerialization_partitionc                 C   s6   | j jD ] }|jtjjjkrtjjj|_q|   | S N)	graphnodestargettorchopsaten_to_copyto	recompile)fx_gnode r   c/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/torch/_functorch/compilers.py_canonicalize    s
    r!   c               	   c   s6   t jd} zd V  W t j|  nt j|  0 d S )NF)r   _CZ_jit_set_autocast_mode)Zold_jit_autocast_flagr   r   r    _disable_jit_autocast(   s    r#   )r   returnc                 C   s0  t   t|  | jjD ]F}|jtjjjkrt	|j
dkrt	|jdkrd|jv rtjjj|_q| jjD ]<}i }|j D ]"\}}t|tjr|j}|||< q|||_qj| j  |   tj| }tj|j tj| }tj|}tdd |D s||  W d   n1 s"0    Y  |S )a  
    Compiles the :attr:`fx_g` with Torchscript compiler.

    .. warning::
        This API is experimental and likely to change.

    Args:
        fx_g(fx.GraphModule): The input Fx graph module to be compiled.

    Returns:
        Torch scripted model.
    r   dtypec                 s   s   | ]}t |tjjV  qd S r   )
isinstancer   Z_subclassesZ
FakeTensor).0tr   r   r    	<genexpr>^       zts_compile.<locals>.<genexpr>N)r#   r   r   r   r   r   r   r   r   lenargskwargsr   itemsr&   devicetypeZlintr   jitscriptr"   Z_jit_pass_remove_mutationfreezeevalZoptimize_for_inferenceany)r   inpsr   Z
new_kwargskvfr   r   r    
ts_compile1   s8    


(r:   Tc                 C   s   t | j t| ||d | S )N)
clear_meta)printcoder   )r   _namer;   r   r   r    _draw_graph_compilec   s    
r@   c                 C   s   t tt| dS )Nr?   )r   r   r@   rA   r   r   r    draw_graph_compilei   s    
rB   c                 C   s   | S )z
    Returns the :attr:`fx_g` Fx graph module as it is. This is a no-op compiler
    and can be used to check accuracy.

    .. warning::
        This API is experimental and likely to change.

    r   r   r>   r   r   r    nopo   s    
rD   c                       s(   e Zd Z fddZ fddZ  ZS )DebugInterpreterc                    s$   t | jg|R  | _t j|  d S r   )r
   modulesymbol_mappingsuperrun)selfr,   	__class__r   r    rI   |   s    zDebugInterpreter.runc           
         s   dd l fddfddfdd  fdd	}t |}d
|jv rt|jd
 \}}t|\}}t|t|ksJ t| dt| ttt|||D ].\}}	t	|	t
jsq|||	fdd q|S )Nr   c                    sB   t | ts| S | jj j}t|jdks:J |t	|S )Nr   )
r&   r   expandr   exprZxreplacerG   r+   Zfree_symbolsint)nir)rJ   sympyr   r    subst_symint   s
    
z/DebugInterpreter.run_node.<locals>.subst_symintc                    s   t  fdd| D S )Nc                 3   s   | ]} |V  qd S r   r   )r'   rP   rS   r   r    r)      r*   zHDebugInterpreter.run_node.<locals>.subst_symint_tuple.<locals>.<genexpr>)tuple)ZnisrT   r   r    subst_symint_tuple   s    z5DebugInterpreter.run_node.<locals>.subst_symint_tuplec                    sT    |   dkrPt| jD ]4} | |||kr | |dkr dS qdS )Nr   r   FT)Znumelrangendimstridesize)abidxrT   r   r    check_significant_strides   s
    *z<DebugInterpreter.run_node.<locals>.check_significant_stridesc              	      s   t |sJ | j|jks6J |  d| j d|j |  | kszJ |  d|   d|   d|   | |}|sJ |  d|   d|   d|  d S )Nz:  != z aka )callabler%   rZ   rY   )nvrvdescZsame_strides)r^   rV   r   r    check   s    **
z(DebugInterpreter.run_node.<locals>.checkvalr_   c                      s   d  dj  S )Nzoutput z where )rG   r   )irJ   r   r    <lambda>   r*   z+DebugInterpreter.run_node.<locals>.<lambda>)rR   rH   run_nodemetapytreeZtree_flattenr+   ziprW   r&   r   Tensor)
rJ   nrd   rQ   Zn_valsZn_specZr_valsZr_specra   rb   rK   )r^   rf   rJ   rS   rV   rR   r    rh      s    
*zDebugInterpreter.run_node)__name__
__module____qualname__rI   rh   __classcell__r   r   rK   r    rE   {   s   rE   c                 C   s
   t | jS )z
    Returns a (slow) interpreter over the FX graph module that also checks
    various debugging properties (e.g., that tracing strides matched real
    strides.)
    )rE   rI   rC   r   r   r    	debug_nop   s    rr   c                 C   s(   t |  tj| }tj| }|S r   )r   r   r1   r2   r3   r4   )r   r>   r9   r   r   r    simple_ts_compile   s    rs   c                 C   s   t | t|dS )N)static_argnums)r   rs   )r9   rt   r   r   r    nnc_jit   s    ru   c                 C   s   t | j | S r   )r<   r=   rC   r   r   r    print_compile   s    
rv   )fnrt   c                 K   sL   t t tt|d}|| t| tjjr8t| fi |S t	| fi |S dS )a  
    Wrapper function over :func:`aot_function` and :func:`aot_module` to perform
    memory efficient fusion. It uses the
    :func:`min_cut_rematerialization_partition` partitioner to perform efficient
    recomputation. It uses NVFuser to compile the generated forward and backward
    graphs.

    .. warning::
        This API is experimental and likely to change.

    Args:
        fn (Union[Callable, nn.Module]): A Python function or a ``nn.Module``
            that takes one ore more arguments. Must return one or more Tensors.
        static_argnums (Optional[Tuple[Int]]): An option tuple of ints to mark
            the arguments of the function as static.
        **kwargs: Any other overrides you want to make to the settings

    Returns:
        Returns a ``Callable``  or ``nn.Module`` that retains the eager behavior
        of the original :attr:`fn`, but whose forward and backward graphs have
        gone through recomputation optimizations, and the graphs have been
        compiled with nvfuser.

    )fw_compilerbw_compilerpartition_fndecompositionsrt   N)
r:   r   default_decompositionsupdater&   r   nnModuler   r   )rw   rt   r-   configr   r   r    memory_efficient_fusion   s    
r   c                 C   sH   |  d tddd |D  d ddlm} |  |  t| |S )NfooaQ  
##############################################################
# To minimize FX graph, copy and paste the below and run it  #
##############################################################

import torch
import torch.fx as fx
from functorch.compile import minifier, check_nvfuser_subprocess, check_nvfuser_correctness_subprocess

inps = c                 S   s   g | ]}|j |jfqS r   )shaper%   )r'   rf   r   r   r    
<listcomp>  r*   z!debug_compile.<locals>.<listcomp>a?  
inps = [torch.ones(shape, dtype=dtype, device='cuda') for (shape, dtype) in inps]
from foo import FxModule
mod = FxModule().cuda()

with torch.jit.fuser("fuser2"):
  # check_nvfuser_subprocess can be replaced with check_nvfuser_correctness_subprocess
  minifier(fx.symbolic_trace(mod), inps, check_nvfuser_subprocess)
r   )FxModule)	to_folderr<   r   r   Zcudar:   )r   r6   r   r   r   r    debug_compile  s    
	r   c                 C   s   g }t | d}t|}g }|D ]}t|dkrD|}|t }nX|\}}}}	}
|	tjtjtj	tj
tjtjtthv rtjdd||	|
d}ntj||	|
d}|| q"W d   n1 s0    Y  |S )zZ
    Return a random input for the given inputs meta generated from _save_fx_default.
    rbr   r   )r%   r/   N)openpickleloadr+   randomZrandr   rO   int32int64boolZuint8floatrandintappend)Zinput_data_pathinputsr9   Zinputs_metari   r0   inputr   rY   r%   r/   r   r   r    
get_inputs1  s.    

*r   c           	         sb   ddl m} fdd fddfdd}fd	d
}fdd}||||||tdS )aO  
    The forward, backward, and joint computation graph will be stored in
    {folder_name}/{current_name}/{current_name}_forward_{graph_index},
    {folder_name}/{current_name}/{current_name}_backward_{graph_index}, and
    {folder_name}/{current_name}/{current_name}_joint_{graph_index} respectively.
    The input shape of the graphs will be stored in the .input files.
    These files can be loaded with pickle,
    and is a list of format (type, shape, stride, dtype, device).
    In the case of type = int or float, it is just (type,).
    For joint graph input, it is a nested list [[],[]]
    where the two inner lists have the same format.
    If dump_example_input is True, example_inputs will be stored in .pt file.
    Since each function might produce multiple graphs,
    the graph_index is used to distinguish difference graphs
    r   )aot_module_simplifiedc                    s   g }t | dkrBt| d trB| | d 7 }| | d 7 }|S | D ]P}t|tksbt|tkrt|t|f qF|t||j| |j	|j
f qF|S )Nr   r   )r+   r&   rU   r0   rO   r   r   r   rY   r%   r/   )r,   
input_metaarg)get_input_metar   r    r   b  s    z(_save_fx_default.<locals>.get_input_metac                    s@  t | jjdkr6ttjd  d| dt d d S t| }|j	t
jj  |  |}tj d  }|st d   | d  d  d| dt 	 t|t d  d  d| dt d  d| dt dd r<t
| d  d  d| dt d  d| dt d d S )	Nr   zNo nodes in graph r>   ./z.inputwbz.pt)r+   r   r   logloggingWARNINGgraph_indexcopydeepcopyZset_codegenr   fxZCodeGenr   ospathexistsmakedirsr   r   dumpr   save)Z
gm_to_saver,   	type_namegmr   ZisExist)current_namedump_example_inputfolder_namer   r   r    graph_saver_helperq  s8    
22z,_save_fx_default.<locals>.graph_saver_helperc                    s    | |d | S )NZforwardr   )r   Zfw_argsr   r   r    graph_saver_forward  s    z-_save_fx_default.<locals>.graph_saver_forwardc                    s    | |d t d7 a | S )NZbackwardr   )r   )r   Zbw_argsr   r   r    graph_saver_backward  s    z._save_fx_default.<locals>.graph_saver_backwardc                    s    | |d t | |S )NZjoint)r   )r   Z
joint_argsr   r   r    graph_saver_joint  s    z+_save_fx_default.<locals>.graph_saver_joint)rx   ry   rz   r{   )Zfunctorch.compiler   r|   )	r   r   r   r   Zexample_inputsr   r   r   r   r   )r   r   r   r   r   r    _save_fx_defaultP  s    "r   Fc                 C   s   da tt| ||S )as  
    Dump the forward, backward, and joint computation graph.
    Example Usage:
    save_fx_func = graph_dumper_aot(current_name, folder_name, dump_example_input = False)
    optimize_ctx = torchdynamo.optimize(
        save_fx_func
    )
    with torch.enable_grad():
        with optimize_ctx:
            result = forward_and_backward_pass(model, example_inputs)
    r   )r   r   r   )r   r   r   r   r   r    graph_dumper_aot  s    r   )T)N)N)F)Xr   r   r   r   r   
contextlibr   	functoolsr   typingr   r   r   r   r   r   Ztorch.fxr   Ztorch.nnr~   Ztorch._decompr	   Z%torch.fx.experimental.symbolic_shapesr
   Zaot_autogradr   r   r   Zcompile_utilsr   Zpartitionersr   r   r   Ztorch.utils._pytreeutilsZ_pytreerj   	getLoggerrn   r   r!   r#   ZGraphModuler:   r@   rB   rD   ZInterpreterrE   rr   rs   ru   r   r   detachZgelu_backwardZleaky_relu_backwardZsigmoid_backwardZthreshold_backwardZhardtanh_backwardZhardsigmoid_backwardZhardswish_backwardZtanh_backwardZsilu_backwardZelu_backwardZcudnn_batch_normZcudnn_batch_norm_backwardZmasked_fillZScalarrl   ZeluZ
leaky_reluZhardtanhZ	hardswishZhardsigmoidZconj_physicalZis_same_sizer|   rv   r   rO   r   r   r   r   r   r   r   r   r   r    <module>   s   

1
3


 
+\