a
    d3                     @   s   d dl Z d dlZd dlmZ ddlmZmZ ddlmZmZ ddlm	Z	 G dd dej
jZdd	ddddddddd

ddZdd Zdd ZdS )    N   )SOLVERSodeint)_check_inputs_flat_to_shape)_mixed_normc                   @   s$   e Zd Zedd Zedd ZdS )OdeintAdjointMethodc                 G   s   || _ || _|
| _|| _|| _|| _|| _|	d u| _t	 h t
||||||||	d}|	d u rx|}| j||g|R   n|\}}| j|||g|R   W d    n1 s0    Y  |S )N)rtolatolmethodoptionsevent_fn)shapesfuncadjoint_rtoladjoint_atoladjoint_methodadjoint_optionst_requires_grad
event_modetorchno_gradr   Zsave_for_backward)ctxr   r   y0tr	   r
   r   r   r   r   r   r   r   r   adjoint_paramsansyevent_t r   b/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/torchdiffeq/_impl/adjoint.pyforward   s     

4zOdeintAdjointMethod.forwardc                    sb  t    | j| j}| j}| j}| j}| j| j}|rt| j	^}}}	 |}
t 
|d d|	dg}|d }n| j	^}} |d }t  t jd|j|jd|d |d g}|dd  D   fdd	}rt jt||j|jd}nd }tt|d ddD ]}rZ|| || }|d|| d}|d  |8  < |||< t|t|||d |d  d||||d
}dd |D }||d  |d< |d  ||d  7  < qr|d |d< |rrt 
|d dt |
dd  g}|d }|dd  }W d    n1 s20    Y  d d ||d d d d d d d d d d g|R S )Nr   r   r   )dtypedevicec                 S   s   g | ]}t |qS r   r   
zeros_like.0paramr   r   r    
<listcomp>B       z0OdeintAdjointMethod.backward.<locals>.<listcomp>c                    s  |d }|d }t   |  }|d} | d}rD| n||}t | dd}t |dd}tdd  D }t jj|| |f  | ddd^}	}
}W d    n1 s0    Y  |	d u rt | n|	}	|
d u rt |n|
}
dd	 t	 |D }|	||
g|R S )
Nr      Tr   c                 s   s   | ]}t |d d V  qdS )r   N)r   
as_stridedr'   r   r   r    	<genexpr>]   r+   zKOdeintAdjointMethod.backward.<locals>.augmented_dynamics.<locals>.<genexpr>)Zallow_unusedZretain_graphc                 S   s&   g | ]\}}|d u rt |n|qS Nr%   )r(   r)   Z	vjp_paramr   r   r    r*   g   s   zLOdeintAdjointMethod.backward.<locals>.augmented_dynamics.<locals>.<listcomp>)
r   Zenable_graddetachZrequires_grad_r-   tupleautogradZgradr&   zip)r   Zy_augr   adj_yZt_	func_eval_tZ_y_paramsZvjp_tZvjp_yZ
vjp_paramsr   r   r   r   r    augmented_dynamicsI   s(    

*z8OdeintAdjointMethod.backward.<locals>.augmented_dynamics)r	   r
   r   r   c                 S   s   g | ]}|d  qS )r   r   )r(   ar   r   r    r*      r+   r,      )r   r   r   r   r   r   r   r   r   Zsaved_tensorscatZreshaper1   zerosr#   r$   extendemptylenrangedotr   Zflipr&   )r   Zgrad_yr   r   r   r   r   r   r   r   r6   Z	aug_stater9   Z	time_vjpsir5   Z	dLd_cur_tr4   
adj_paramsr   r8   r    backward$   sV    
"'&,zOdeintAdjointMethod.backwardN)__name__
__module____qualname__staticmethodr!   rE   r   r   r   r    r   	   s   
r   gHz>g&.>)
r	   r
   r   r   r   r   r   r   r   r   c       
         C   s  |d u rt | tjstd|d u r(|}|	d u r4|}	|
d u r@|}
|
|kr`|d ur`|d u r`td|d u r|d urdd | D ni }n| }|d u rtt| }nt|}t|}tdd |D }t||krd|v rt	|d rt
d t| |||||||t	\
}} }}}}}}}}|d }t||| tj|| |||||||||	|
||jg|R  }|d u rp|}n|\}}||}|r| }|d urt|t|f|}|d u r|S ||fS d S )	Nzfunc must be an instance of nn.Module to specify the adjoint parameters; alternatively they can be specified explicitly via the `adjoint_params` argument. If there are no parameters then it is allowable to set `adjoint_params=()`.zIf `adjoint_method != method` then we cannot infer `adjoint_options` from `options`. So as `options` has been passed then `adjoint_options` must be passed as well.c                 S   s   i | ]\}}|d kr||qS )normr   r(   kvr   r   r    
<dictcomp>   r+   z"odeint_adjoint.<locals>.<dictcomp>c                 s   s   | ]}|j r|V  qd S r/   )requires_grad)r(   pr   r   r    r.      r+   z!odeint_adjoint.<locals>.<genexpr>rJ   zAn adjoint parameter was passed without requiring gradient. For efficiency this will be excluded from the adjoint pass, and will not appear as a tensor in the adjoint norm.)
isinstancennModule
ValueErroritemscopyr1   find_parametersr@   callablewarningswarnr   r   handle_adjoint_norm_r   applyrO   tor   )r   r   r   r	   r
   r   r   r   r   r   r   r   r   Zoldlen_r   Zdecreasing_time
state_normr   Zsolutionr   r   r   r    odeint_adjoint   sP     
,



r_   c                 C   sN   t | tjsJ t| ddr>dd }| j|d}dd |D S t|  S d S )NZ_is_replicaFc                 S   s   dd | j  D }|S )Nc                 S   s(   g | ] \}}t |r|jr||fqS r   )r   Z	is_tensorrO   rK   r   r   r    r*      r+   zCfind_parameters.<locals>.find_tensor_attributes.<locals>.<listcomp>)__dict__rU   )moduleZtuplesr   r   r    find_tensor_attributes   s    z/find_parameters.<locals>.find_tensor_attributes)Zget_members_fnc                 S   s   g | ]\}}|qS r   r   )r(   _r)   r   r   r    r*      r+   z#find_parameters.<locals>.<listcomp>)rQ   rR   rS   getattrZ_named_memberslist
parameters)ra   rb   genr   r   r    rW      s    rW   c                    s   fdd}d| vr|| d< nfz| d  W n t yD   || d< Y n@0  dkrdfdd}|| d< n du rnn fdd	}|| d< dS )
zJIn-place modifies the adjoint options to choose or wrap the norm function.c                    s*   | ^}}}}t |  | |t|S r/   )maxabsr   Ztensor_tupler   r   r4   rD   r^   r   r    default_adjoint_norm   s    z2handle_adjoint_norm_.<locals>.default_adjoint_normrJ   Zseminormc                    s$   | ^}}}}t |  | |S r/   )rh   ri   rj   rk   r   r    adjoint_seminorm  s    z.handle_adjoint_norm_.<locals>.adjoint_seminormNc                    s<   | ^}}}}t |d}t |d} |g|||R S )Nr   )r   rj   )adjoint_normr   r   r    _adjoint_norm  s    z+handle_adjoint_norm_.<locals>._adjoint_norm)KeyError)r   r   r^   rl   rm   ro   r   )rn   r   r^   r    r[      s    

r[   )rY   r   Ztorch.nnrR   r   r   miscr   r   r   r2   ZFunctionr   r_   rW   r[   r   r   r   r    <module>   s    

F