a
    d                     @   s   d dl Z 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mZ ddl	m
Z
 eeZe
edddddZe jed	d
Ze jedd
Zdd Ze ddd ZdS )    N   )device_from_inputsfake_tensor_unsupported)register_backend N  )	schedulertrialsc             
      s:  dd l ddl m} ddlm} tj| |}t|}dd t|D }|j	
||\}	}
|jdkr||j}j }nd}jt }|d u rtjdd }|dkrdd	l m} t }tj|s||	d
 |
|\}}|D ]}t|j qtd t|dkr|||}tj|s|dks>J |j|| |gdd}z|!| W n. t"y   tj|rt#|  Y n0 |$|V j%j&dddid  |j'|	||
d}W d    n1 s0    Y  W d    n1 s0    Y  n|dkrddl m(} t) l}|jdkrVjt  d|j*j+dd }|j,j-|	||dd|
dd}|j,j.||	||
d}W d    n1 s0    Y  nZ|dks|sj%j&dd  |j'|	||
d}W d    n1 s0    Y  nt/d |0|d | d!d" fd#d$ fd%d&}|S )'Nr   )relay)graph_executorc                 S   s    g | ]\}}d | |j fqS )inp_)shape).0idxi r   c/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/torch/_dynamo/backends/tvm.py
<listcomp>       ztvm.<locals>.<listcomp>cudaZTVM_SCHEDULERauto_scheduler)r   mainzNo tasksi  )Znum_measure_trialsZmeasure_callbacksZearly_stopping   z relay.backend.use_auto_schedulerT)	opt_levelconfig)targetparamsmeta_schedule)r   z --num-cores F)logicalr   @   Zevolutionary)modr   work_dirZmax_trials_globalZnum_trials_per_iterr   Zstrategy)databaser   r   r   default
   )r   zThis tuning option is invalid/not implemented for torchdynamo's TVM-related backend. There are three available options: default, auto_scheduler and meta_schedule.c                 S   s*   | j dkrt|  S tjj|  S )z8A helper function to transfer a NDArray to torch.tensor.bool)dtypetorchZ
from_numpynumpyutilsZdlpackfrom_dlpackZ	to_dlpack)Z	nd_tensorr   r   r   to_torch_tensorl   s    
ztvm.<locals>.to_torch_tensorc                    s,   | j tjkr  j|   S  j| S )z8A helper function to transfer a torch.tensor to NDArray.)r%   r&   r$   Zndarraycpur'   r)   )Ztorch_tensor)tvmr   r   to_tvm_tensoru   s    ztvm.<locals>.to_tvm_tensorc                     sv   dd | D }t |dD ]8\}}| dkr|jr:| } d| | q    fddt  D S )Nc                 S   s   g | ]}|  qS r   )
contiguous)r   ar   r   r   r   ~   r   z)tvm.<locals>.exec_tvm.<locals>.<listcomp>r   r   c                    s   g | ]}  |qS r   )Z
get_output)r   r   )mr*   r   r   r      r   )	enumerateZdimZrequires_graddetachZ	set_inputrunrangeZget_num_outputs)Zi_argsargsr   arg)r1   r*   r.   r   r   exec_tvm}   s    ztvm.<locals>.exec_tvm)1r-   r	   Ztvm.contribr
   r&   Zjittracer   r2   ZfrontendZfrom_pytorchtyper   indexr   r,   ZTargetllvm_targetosenvirongetr   tempfileNamedTemporaryFilepathexistsZextract_tasksprintZcompute_daglenZTaskSchedulerZTuningOptionsZRecordToFileZtune	ExceptionunlinkZApplyHistoryBestZ	transformZPassContextbuildr   TemporaryDirectoryr(   	cpu_countZrelay_integrationZ
tune_relayZcompile_relayNotImplementedErrorZGraphModule)ZgmZexample_inputsr   r   r	   r
   Zjit_modZdeviceZ
shape_listr   r   devr   r   Zlog_fileZtasksZtask_weightsZtaskZtunerZtune_optionlibmsr    r!   r8   r   )r1   r*   r.   r-   r   r-      s    





R

	(2	r-   r   )r   r   c                   C   s*   zt d W dS  ty$   Y dS 0 d S )Nr-   TF)	importlibimport_moduleImportErrorr   r   r   r   has_tvm   s
    
rR   c                   C   s   dt d v rdS dS )NZavx512z/proc/cpuinfozllvm -mcpu=skylake-avx512zllvm -mcpu=core-avx2)openreadr   r   r   r   r<      s    r<   )	functoolsrO   loggingr=   r@   r&   commonr   r   registryr   	getLogger__name__logr-   partialZtvm_meta_scheduleZtvm_auto_schedulerrR   	lru_cacher<   r   r   r   r   <module>   s    
|