a
    d-                     @   sR  d dl mZ d dlZd dlZd dlZd dlZddlmZ dd Z	dd Z
d	d
 Zdd Zdd Zd7ddZdd Ze dd Zdd Zdd Zdd Zdd Zdd Zdd  ZG d!d" d"ejjZG d#d$ d$ejjZG d%d& d&ejjZG d'd( d(eZG d)d* d*ejjZd+d, ZG d-d. d.ej j!Z"d/d0 Z#d1d2 Z$d3d4 Z%d5d6 Z&dS )8    )EnumN   )combine_event_functionsc                 C   s&   t |dkr"td| jj| d S )Nr   z{}: Unexpected arguments {})lenwarningswarnformat	__class____name__)ZsolverZunused_kwargs r   _/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/torchdiffeq/_impl/misc.py_handle_unused_kwargs	   s    r   c                 C   s   |   S N)maxtensorr   r   r   
_linf_norm   s    r   c                 C   s   |  d  S )N   )powmeansqrtr   r   r   r   	_rms_norm   s    r   c                 C   s   dS )N        r   r   r   r   r   
_zero_norm   s    r   c                 C   s"   t | dkrdS tdd | D S )Nr   r   c                 S   s   g | ]}t |qS r   )r   ).0r   r   r   r   
<listcomp>       z_mixed_norm.<locals>.<listcomp>)r   r   )Ztensor_tupler   r   r   _mixed_norm   s    r   c                 C   s  |j }|j}	|j }
||}|du r.| ||}|t||  }||| }||| }|dk sh|dk rztjd||	d}nd| | }|||  }| || |}||| | | }|dkr|dkrttjd||	d|d }ndt|| dt|d	   }td
| ||
S )a  Empirically select a good initial step.

    The algorithm is described in [1]_.

    References
    ----------
    .. [1] E. Hairer, S. P. Norsett G. Wanner, "Solving Ordinary Differential
           Equations I: Nonstiff Problems", Sec. II.4, 2nd edition.
    Ngh㈵>gư>dtypedeviceg{Gz?gV瞯<gMbP?      ?r   d   )	r   r    totorchabsr   r   floatmin)funct0y0orderrtolatolnormZf0r   r    Zt_dtypeZscaleZd0d1Zh0y1f1Zd2h1r   r   r   _select_initial_step    s&    

r3   c                 C   s(   ||t | |   }|| | S r   )r$   r   r%   )Zerror_estimater,   r-   r*   r0   r.   Z	error_tolr   r   r   _compute_error_ratioJ   s    r4   c                 C   sr   |dkr| | S |dk r,t jd| j| jd}|| }t j|| j| jd }t |t |||  |}| | S )z-Calculate the optimal size for the next step.r   r   r   r   )	r$   Zonesr   r    Ztype_asr   Z
reciprocalr'   r   )Z	last_stepZerror_ratioZsafetyZifactorZdfactorr+   exponentZfactorr   r   r   _optimal_step_sizeO   s    
r6   c                 C   s   | dd  | d d k   S )Nr   )all)tr   r   r   _decreasing\   s    r:   c                 C   s   |  dksJ d| d S )Nr   {} must be one dimensional)
ndimensionr   namer9   r   r   r   _assert_one_dimensional`   s    r?   c                 C   s.   |dd  |d d k  s*J d| d S )Nr   r7   ,{} must be strictly increasing or decreasing)r8   r   r=   r   r   r   _assert_increasingd   s    rA   c                 C   s"   t |std| | d S )Nz0`{}` must be a floating point Tensor but is a {})r$   Zis_floating_point	TypeErrorr   typer=   r   r   r   _assert_floatingh   s    
rD   c                 C   sh   zt | W n ty"   | Y S 0 t|}t|t|ksJJ d| dd t||D }t|S )Nz?If using tupled {} it must have the same length as the tuple y0c                 S   s$   g | ]\}}t || qS r   )r$   Z	as_tensorexpandnumel)r   Ztol_shaper   r   r   r   t   r   z_tuple_tol.<locals>.<listcomp>)iterrB   tupler   r   zipr$   cat)r>   Ztolshapesr   r   r   
_tuple_tolm   s    
rM   c                 C   sP   g }d}|D ]:}||   }|| d||f g ||R  |}qt|S )Nr   .)rF   appendviewrI   )r   lengthrL   Ztensor_listtotalrG   Z
next_totalr   r   r   _flat_to_shapex   s    &rR   c                       s$   e Zd Z fddZdd Z  ZS )
_TupleFuncc                    s   t t|   || _|| _d S r   )superrS   __init__	base_funcrL   selfrV   rL   r	   r   r   rU      s    z_TupleFunc.__init__c                 C   s*   |  |t|d| j}tdd |D S )Nr   c                 S   s   g | ]}| d qS r7   Zreshape)r   Zf_r   r   r   r      r   z&_TupleFunc.forward.<locals>.<listcomp>)rV   rR   rL   r$   rK   )rX   r9   yfr   r   r   forward   s    z_TupleFunc.forwardr
   
__module____qualname__rU   r^   __classcell__r   r   rY   r   rS      s   rS   c                       s$   e Zd Z fddZdd Z  ZS )_TupleInputOnlyFuncc                    s   t t|   || _|| _d S r   )rT   rc   rU   rV   rL   rW   rY   r   r   rU      s    z_TupleInputOnlyFunc.__init__c                 C   s   |  |t|d| jS Nr   )rV   rR   rL   rX   r9   r\   r   r   r   r^      s    z_TupleInputOnlyFunc.forwardr_   r   r   rY   r   rc      s   rc   c                       s&   e Zd Zd fdd	Zdd Z  ZS )_ReverseFuncr!   c                    s   t t|   || _|| _d S r   )rT   rf   rU   rV   mul)rX   rV   rg   rY   r   r   rU      s    z_ReverseFunc.__init__c                 C   s   | j | | | S r   )rg   rV   re   r   r   r   r^      s    z_ReverseFunc.forward)r!   r_   r   r   rY   r   rf      s   rf   c                   @   s   e Zd ZdZdZdZdS )Perturbr   r   r   N)r
   r`   ra   NONEPREVNEXTr   r   r   r   rh      s   rh   c                       s,   e Zd Z fddZejdddZ  ZS )_PerturbFuncc                    s   t t|   || _d S r   )rT   rl   rU   rV   )rX   rV   rY   r   r   rU      s    z_PerturbFunc.__init__)perturbc                C   s^   t |tsJ d||j}|tju r8t||d }n|tju rRt||d }n | ||S )Nz-perturb argument must be of type Perturb enumr   )
isinstancerh   r#   r   rk   
_nextafterrj   rV   )rX   r9   r\   rm   r   r   r   r^      s    

z_PerturbFunc.forward)r
   r`   ra   rU   rh   ri   r^   rb   r   r   rY   r   rl      s   rl   c	              
      s  |d ur8t |dkr(tdt | dt||d |}d t|tj }	|	rt|ts`J ddd |D td|}td	|}td
d |D }t	| } |d urt
|}td| |d u ri }n| }|d u rd}||vrtd|dd|  d |	rDd|v r(|d ntfdd}
|
|d< nd|v rPnt|d< td|d d}t |dkr|d |d krd}|r| }t| dd} |d urt|}z|d  W n ty   Y n0  fdd|d< t|d t|d td| t|r"|jr"J dt|r>|jr>J d |j|jkrbtd! ||j}t| } | ||||||||f
S )"Nr   zCWe require len(t) == 2 when in event handling mode, but got len(t)=.r   z+y0 must be either a torch.Tensor or a tuplec                 S   s   g | ]
}|j qS r   )rG   r   Zy0_r   r   r   r      r   z!_check_inputs.<locals>.<listcomp>r,   r-   c                 S   s   g | ]}| d qS rZ   r[   rq   r   r   r   r      r   r*   Zdopri5z&Invalid method "{}". Must be one of {}z{"z", "z"}.r.   c                    s   t | d} |S rd   )rR   )r   r\   )r.   rL   r   r   _norm   s    z_check_inputs.<locals>._normr9   TFr   g      )rg   Zgrid_constructorc                    s    | ||  S r   r   )r(   r*   r9   )_grid_constructorr   r   <lambda>  r   z_check_inputs.<locals>.<lambda>Zstep_tZjump_tzrtol cannot require gradientzatol cannot require gradientz9t is not on the same device as y0. Coercing to y0.device.)r   
ValueErrorr   rn   r$   TensorrI   rM   rK   rS   rc   rD   copyr   joinkeysr   r   _check_timelikerf   KeyError_flip_optionrA   Z	is_tensorrequires_gradr    r   r   r#   rl   )r(   r*   r9   r,   r-   methodoptionsZevent_fnZSOLVERSZis_tuplerr   Zt_is_reversedr   )rs   r.   rL   r   _check_inputs   sx    







 




r   c                   @   s$   e Zd Zedd Zedd ZdS )_StitchGradientc                 C   s   |S r   r   )ctxx1outr   r   r   r^   3  s    z_StitchGradient.forwardc                 C   s   |d fS r   r   )r   Zgrad_outr   r   r   backward7  s    z_StitchGradient.backwardN)r
   r`   ra   staticmethodr^   r   r   r   r   r   r   2  s   
r   c                 C   sV   t  2 tt dr"t | |}n
t| |}W d    n1 s@0    Y  t| |S )N	nextafter)r$   no_gradhasattrr   np_nextafterr   apply)r   x2r   r   r   r   ro   <  s
    

(ro   c                 C   sF   t d |    }|   }tt||	| }|S )Nztorch.nextafter is only available in PyTorch 1.7 or newer.Falling back to numpy.nextafter. Upgrade PyTorch to remove this warning.)
r   r   detachcpunumpyr$   r   npr   r#   )r   r   Zx1_npZx2_npr   r   r   r   r   E  s
    
r   c                 C   s   t |tjsJ d| t| | | dks>J d| |sV|jrVJ d| |dd  |d d k}| s|  sJ d| d S )Nz{} must be a torch.Tensorr   r;   z{} cannot require gradientr7   r@   )rn   r$   rv   r   rD   r<   r}   r8   )r>   ZtimelikeZcan_graddiffr   r   r   rz   N  s    
rz   c                 C   s:   z| | }W n t y   Y n0 t|tjr6| | |< d S r   )r{   rn   r$   rv   )r   Zoption_nameZoption_valuer   r   r   r|   X  s    r|   )N)'enumr   mathr   r   r$   r   Zevent_handlingr   r   r   r   r   r   r3   r4   r   r6   r:   r?   rA   rD   rM   rR   nnModulerS   rc   rf   rh   rl   r   ZautogradZFunctionr   ro   r   rz   r|   r   r   r   r   <module>   s>   
*


r
		
