a
    d                  
   @   s   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m	Z	 G dd dej
ZG d	d
 d
ej
ZG dd dej
Zdeee ee eeedf  ee j ee	ee jf  dddZdS )    N   )brownian_base)brownian_interval   )OptionalScalarTensorTupleUnionc                       s^   e Zd Z fddZdddZdd Zed	d
 Zedd Zedd Z	edd Z
  ZS )ReverseBrownianc                    s   t t|   || _d S N)superr   __init__base_brownian)selfr   	__class__ c/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/torchsde/_brownian/derived.pyr      s    zReverseBrownian.__init__NFc                 C   s   | j | | ||dS N)return_Ureturn_A)r   )r   tatbr   r   r   r   r   __call__   s    zReverseBrownian.__call__c                 C   s   | j j d| j dS )Nz(base_brownian=))r   __name__r   r   r   r   r   __repr__    s    zReverseBrownian.__repr__c                 C   s   | j jS r   )r   dtyper   r   r   r   r   #   s    zReverseBrownian.dtypec                 C   s   | j jS r   )r   devicer   r   r   r   r    '   s    zReverseBrownian.devicec                 C   s   | j jS r   )r   shaper   r   r   r   r!   +   s    zReverseBrownian.shapec                 C   s   | j jS r   )r   levy_area_approximationr   r   r   r   r"   /   s    z'ReverseBrownian.levy_area_approximation)NFF)r   
__module____qualname__r   r   r   propertyr   r    r!   r"   __classcell__r   r   r   r   r      s   



r   c                       sn   e Zd ZdZdeeed fddZddd	Zd
d Z	e
dd Ze
dd Ze
dd Ze
dd Z  ZS )BrownianPatha  Brownian path, storing every computed value.

    Useful for speed, when memory isn't a concern.

    To use:
    >>> bm = BrownianPath(t0=0.0, w0=torch.zeros(4, 1))
    >>> bm(0., 0.5)
    tensor([[ 0.0733],
            [-0.5692],
            [ 0.1872],
            [-0.3889]])
       )t0w0window_sizec                    s>   |d }|| _ tj|||j|j|jdd| _tt| 	  dS )zInitialize Brownian path.
        Arguments:
            t0: Initial time.
            w0: Initial state.
            window_size: Unused; deprecated.
        r   N)r)   t1sizer   r    Z
cache_size)
_w0r   BrownianIntervalr!   r   r    	_intervalr   r'   r   )r   r)   r*   r+   r,   r   r   r   r   B   s    zBrownianPath.__init__NFc                 C   s0   | j ||||d}|d u r,|s,|s,|| j }|S r   r0   r.   r   tr   r   r   outr   r   r   r   O   s    
zBrownianPath.__call__c                 C   s   | j j d| j dS Nz
(interval=r   r   r   r0   r   r   r   r   r   V   s    zBrownianPath.__repr__c                 C   s   | j jS r   r0   r   r   r   r   r   r   Y   s    zBrownianPath.dtypec                 C   s   | j jS r   r0   r    r   r   r   r   r    ]   s    zBrownianPath.devicec                 C   s   | j jS r   r0   r!   r   r   r   r   r!   a   s    zBrownianPath.shapec                 C   s   | j jS r   r0   r"   r   r   r   r   r"   e   s    z$BrownianPath.levy_area_approximation)r(   )NFF)r   r#   r$   __doc__r   r   intr   r   r   r%   r   r    r!   r"   r&   r   r   r   r   r'   4   s   



r'   c                       s   e Zd ZdZdeeee ee ee eeeee d	 fddZ	dd
dZ
dd Zedd Zedd Zedd Zedd Z  ZS )BrownianTreea  Brownian tree with fixed entropy.

    Useful when the map from entropy -> Brownian motion shouldn't depend on the
    locations and order of the query points. (As the usual BrownianInterval
    does - note that BrownianTree is slower as a result though.)

    To use:
    >>> bm = BrownianTree(t0=0.0, w0=torch.zeros(4, 1))
    >>> bm(0., 0.5)
    tensor([[ 0.0733],
            [-0.5692],
            [ 0.1872],
            [-0.3889]], device='cuda:0')
    Nư>   	   )	r)   r*   r,   w1entropytol	pool_sizecache_depthsafetyc
                    sd   |du r|d }|du rd}
n|| }
|| _ tj|||j|j|j|||d|
d
| _tt| 	  dS )a  Initialize the Brownian tree.

        The random value generation process exploits the parallel random number paradigm and uses
        `numpy.random.SeedSequence`. The default generator is PCG64 (used by `default_rng`).

        Arguments:
            t0: Initial time.
            w0: Initial state.
            t1: Terminal time.
            w1: Terminal state.
            entropy: Global seed, defaults to `None` for random entropy.
            tol: Error tolerance before the binary search is terminated; the search depth ~ log2(tol).
            pool_size: Size of the pooled entropy. This parameter affects the query speed significantly.
            cache_depth: Unused; deprecated.
            safety: Unused; deprecated.
        Nr   T)
r)   r,   r-   r   r    rB   rC   rD   Zhalfway_treeW)
r.   r   r/   r!   r   r    r0   r   r=   r   )r   r)   r*   r,   rA   rB   rC   rD   rE   rF   rG   r   r   r   r   z   s$    
zBrownianTree.__init__Fc                 C   s0   | j ||||d}|d u r,|s,|s,|| j }|S r   r1   r2   r   r   r   r      s    
zBrownianTree.__call__c                 C   s   | j j d| j dS r5   r6   r   r   r   r   r      s    zBrownianTree.__repr__c                 C   s   | j jS r   r7   r   r   r   r   r      s    zBrownianTree.dtypec                 C   s   | j jS r   r8   r   r   r   r   r       s    zBrownianTree.devicec                 C   s   | j jS r   r9   r   r   r   r   r!      s    zBrownianTree.shapec                 C   s   | j jS r   r:   r   r   r   r   r"      s    z$BrownianTree.levy_area_approximation)NNNr>   r?   r@   N)NFF)r   r#   r$   r;   r   r   r   r<   floatr   r   r   r%   r   r    r!   r"   r&   r   r   r   r   r=   j   s8          -



r=                 ?.)yr)   r,   r-   r   r    c                 K   sR   |du r| j n|}|du r | jn|}|du r2| jn|}tjf |||||d|S )zZReturns a BrownianInterval object with the same size, device, and dtype as a given tensor.N)r)   r,   r-   r   r    )r!   r   r    r   r/   )rK   r)   r,   r-   r   r    kwargsr   r   r   brownian_interval_like   s    rM   )rI   rJ   NNN)Ztorch r   r   typesr   r   r   r	   r
   ZBaseBrownianr   r'   r=   r<   r   strr    rM   r   r   r   r   <module>   s$   6Y     