a
    d                     @   sj   d dl Z d dlZddlmZ ddlmZ G dd de jdZG dd	 d	ee jdZG d
d de jdZ	dS )    N   )
find_event)_handle_unused_kwargsc                   @   s2   e Zd Zdd Zdd Zejdd Zdd Zd	S )
AdaptiveStepsizeODESolverc                 K   s"   t | | ~|| _|| _|| _d S N)r   y0dtypenorm)selfr   r   r	   unused_kwargs r   b/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/torchdiffeq/_impl/solvers.py__init__   s
    
z"AdaptiveStepsizeODESolver.__init__c                 C   s   d S r   r   )r
   tr   r   r   _before_integrate   s    z+AdaptiveStepsizeODESolver._before_integratec                 C   s   t d S r   NotImplementedError)r
   Znext_tr   r   r   _advance   s    z"AdaptiveStepsizeODESolver._advancec                 C   st   t jt|g| jjR | jj| jjd}| j|d< || j}| | t	dt|D ]}| 
|| ||< qX|S )Nr   devicer   r   )torchemptylenr   shaper   r   tor   ranger   )r
   r   solutionir   r   r   	integrate   s    *

z#AdaptiveStepsizeODESolver.integrateN)	__name__
__module____qualname__r   r   abcabstractmethodr   r   r   r   r   r   r      s
   	
r   )	metaclassc                   @   s"   e Zd Zejdd Zdd ZdS )AdaptiveStepsizeEventODESolverc                 C   s   t d S r   r   )r
   event_fnr   r   r   _advance_until_event$   s    z3AdaptiveStepsizeEventODESolver._advance_until_eventc                 C   sL   | | jj| j}| |d | |\}}tj| j|gdd}||fS )Nr   Zdim)	r   r   r   r   r   Zreshaper'   r   stack)r
   t0r&   
event_timey1r   r   r   r   integrate_until_event(   s
    z4AdaptiveStepsizeEventODESolver.integrate_until_eventN)r   r    r!   r"   r#   r'   r.   r   r   r   r   r%   "   s   
r%   c                   @   sZ   e Zd ZU eed< dddZedd Zej	d	d
 Z
dd Zdd Zdd Zdd ZdS )FixedGridODESolverorderNlinearFc                 K   s   | d| _| dd  | dd  t| | ~|| _|| _|j| _|j| _|| _|| _|| _	|d u r|d u rzdd | _
q|| _
n|d u r| || _
ntdd S )NatolZrtolr	   c                 S   s   |S r   r   )fr   r   r   r   r   <lambda>D       z-FixedGridODESolver.__init__.<locals>.<lambda>z@step_size and grid_constructor are mutually exclusive arguments.)popr2   r   funcr   r   r   	step_sizeinterpperturbgrid_constructor _grid_constructor_from_step_size
ValueError)r
   r7   r   r8   r;   r9   r:   r   r   r   r   r   3   s&    
zFixedGridODESolver.__init__c                    s    fdd}|S )Nc                    sX   |d }|d }t ||   d  }t jd||j|jd  | }|d |d< |S )Nr   r(   r   r   )r   ceilitemZaranger   r   )r7   r   r   
start_timeZend_timeZnitersZt_inferr8   r   r   _grid_constructorO   s    zNFixedGridODESolver._grid_constructor_from_step_size.<locals>._grid_constructorr   )r8   rB   r   rA   r   r<   M   s    	z3FixedGridODESolver._grid_constructor_from_step_sizec                 C   s   d S r   r   )r
   r7   r+   dtt1r   r   r   r   
_step_funcZ   s    zFixedGridODESolver._step_funcc                 C   s^  |  | j| j|}|d |d kr2|d |d ks6J tjt|g| jjR | jj| jjd}| j|d< d}| j}t	|d d |dd  D ]\}}|| }| 
| j||||\}	}
||	 }|t|k rT||| krT| jdkr| |||||| ||< nH| jdkr:| ||}| |||
||||| ||< ntd| j |d7 }q|}q|S )Nr   r(   r   r   r1   cubicUnknown interpolation method )r;   r7   r   r   r   r   r   r   r   ziprE   r9   _linear_interp_cubic_hermite_interpr=   )r
   r   Z	time_gridr   jr   r+   rD   rC   dyf0r-   f1r   r   r   r   ^   s(    $*
" 
zFixedGridODESolver.integratec                    sN  j d usJ djjj }t|}d}d}|d7 }| j|\} | t|}||krjdkrfdd}	n@jdkr܈ fd	d}	ntd
j t	|	||t
j\}
q2n
 ||krBtd| dqBtjjgdd}|
|fS )Nz_Event handling for fixed step solvers currently requires `step_size` to be provided in options.i N  r   r   r1   c                    s     | S r   )rI   r   )r
   r+   rD   r   r-   r   r   r4      r5   z:FixedGridODESolver.integrate_until_event.<locals>.<lambda>rF   c              	      s     | S r   )rJ   rO   rM   rN   r
   r+   rD   r   r-   r   r   r4      r5   rG   z%Reached maximum number of iterations .r)   )r8   Ztype_asr   r   signrE   r7   r9   r=   r   floatr2   RuntimeErrorr*   )r
   r+   r&   rC   Zsign0Zmax_itrsitrrL   Zsign1Z	interp_fnr,   r   r   rP   r   r.   y   s4    



z(FixedGridODESolver.integrate_until_eventc                 C   s   || ||  }dd|  d|  d|  }	|d|  d|  }
|| dd|   }|| |d  }|| }|	| |
| |  ||  || |  S )Nr         r   )r
   r+   r   rM   rD   r-   rN   r   hZh00Zh10Zh01Zh11rC   r   r   r   rJ      s    z(FixedGridODESolver._cubic_hermite_interpc                 C   s8   ||kr|S ||kr|S || ||  }||||   S r   r   )r
   r+   rD   r   r-   r   Zsloper   r   r   rI      s    z!FixedGridODESolver._linear_interp)NNr1   F)r   r    r!   int__annotations__r   staticmethodr<   r"   r#   rE   r   r.   rJ   rI   r   r   r   r   r/   0   s   



$	r/   )
r"   r   Zevent_handlingr   miscr   ABCMetar   r%   r/   r   r   r   r   <module>   s   