a
    þd>z  ã                   @   sf  U d dl Z d dlZd dlmZ d dlmZmZmZmZ d dl	Z	d dl	m
Z
 ddlmZ ddlmZ ddlmZ dd	lmZ dd
lmZ ddlmZ ddlmZ ddlmZ ddlmZ ddlmZ ddlm Z  ddl!m"Z" ddl#m$Z$ ddl%m&Z& ddl'm(Z( ddl)m*Z* ddl+m,Z,m-Z-m.Z. ddl/m0Z0m1Z1 ddl2m3Z3 ddl4m5Z5 ddl6m7Z7 ddl8m9Z9 ddl:m;Z; ddl<m=Z= ddl>m?Z?m@ZA i ZBi ZCeeeef ef eDd< d d!gZEd"d „ ZFeG d#d$„ d$ƒƒZGd%d&„ ZHd'd(„ ZId)d*„ ZJd+d,„ ZKeee	jLd-œd.d!„ZMeFeeƒd/d0„ ƒZNeFeeƒd1d2„ ƒZOeFeeƒd3d4„ ƒZPeFeeƒd5d6„ ƒZQeFeeƒd7d8„ ƒZReFeeƒd9d:„ ƒZSeFeeƒd;d<„ ƒZTeFeeƒd=d>„ ƒZUeFe e ƒd?d@„ ƒZVeFe$e$ƒdAdB„ ƒZWeFe"e"ƒdCdD„ ƒZXeFe&e&ƒdEdF„ ƒZYeFe*e*ƒdGdH„ ƒZZeFe,e,ƒdIdJ„ ƒZ[eFe0e,ƒdKdL„ ƒZ\eFe,e0ƒdMdN„ ƒZ]eFe0e0ƒdOdP„ ƒZ^eFe3e3ƒdQdR„ ƒZ_eFe5e5ƒdSdT„ ƒZ`eFe7e7ƒdUdV„ ƒZaeFe9e9ƒdWdX„ ƒZbeFe;e;ƒdYdZ„ ƒZceFe=e=ƒd[d\„ ƒZdeFee9ƒd]d^„ ƒZeeFeeƒd_d`„ ƒZfeFee7ƒdadb„ ƒZgeFeeƒdcdd„ ƒZheFee ƒdedf„ ƒZieFee3ƒdgdh„ ƒZjeFee=ƒdidj„ ƒZkeFee7ƒdkdl„ ƒZleFeeƒdmdn„ ƒZmeFee3ƒdodp„ ƒZneFee=ƒdqdr„ ƒZoeFeeƒeFeeƒeFee7ƒeFee=ƒdsdt„ ƒƒƒƒZpeFee ƒdudv„ ƒZqeFee$ƒdwdx„ ƒZreFee3ƒdydz„ ƒZseFe eƒeFe eƒeFe e7ƒeFe e=ƒd{d|„ ƒƒƒƒZteFe eƒd}d~„ ƒZueFe e$ƒdd€„ ƒZveFe e3ƒdd‚„ ƒZweFe$eƒeFe$eƒeFe$eƒeFe$e ƒeFe$e7ƒeFe$e=ƒdƒd„„ ƒƒƒƒƒƒZxeFe$e3ƒd…d†„ ƒZyeFe*eƒeFe*eƒeFe*eƒeFe*e ƒeFe*e7ƒeFe*e=ƒd‡dˆ„ ƒƒƒƒƒƒZzeFe*e3ƒd‰dŠ„ ƒZ{eFe3eƒeFe3eƒeFe3eƒeFe3e ƒeFe3e7ƒeFe3e=ƒd‹dŒ„ ƒƒƒƒƒƒZ|eFe3e$ƒddŽ„ ƒZ}eFe3e*ƒdd„ ƒZ~eFe7eƒeFe7eƒeFe7e=ƒd‘d’„ ƒƒƒZeFe7eƒd“d”„ ƒZ€eFe7e ƒd•d–„ ƒZeFe7e3ƒd—d˜„ ƒZ‚eFe9eƒeFe9eƒd™dš„ ƒƒZƒeFe=eƒd›dœ„ ƒZ„eFe=eƒddž„ ƒZ…eFe=eƒdŸd „ ƒZ†eFe=e ƒd¡d¢„ ƒZ‡eFe=e$ƒd£d¤„ ƒZˆeFe=e3ƒd¥d¦„ ƒZ‰eFe=e7ƒd§d¨„ ƒZŠeFe(e(ƒd©dª„ ƒZ‹eFeeƒd«d¬„ ƒZŒd­d®„ ZdS )¯é    N)Útotal_ordering)ÚTypeÚDictÚCallableÚTuple)Úinfé   )Ú	Bernoulli)ÚBeta)ÚBinomial)ÚCategorical)ÚCauchy)ÚContinuousBernoulli)Ú	Dirichlet)ÚDistribution)ÚExponential)ÚExponentialFamily)ÚGamma)Ú	Geometric)ÚGumbel)Ú
HalfNormal)ÚIndependent)ÚLaplace)ÚLowRankMultivariateNormalÚ_batch_lowrank_logdetÚ_batch_lowrank_mahalanobis)ÚMultivariateNormalÚ_batch_mahalanobis)ÚNormal)ÚOneHotCategorical)ÚPareto)ÚPoisson)ÚTransformedDistribution)ÚUniform)Ú_sum_rightmostÚeuler_constantÚ_KL_MEMOIZEÚregister_klÚkl_divergencec                    sV   t ˆ tƒs"tˆ tƒr"td ˆ ¡ƒ‚t ˆtƒsDtˆtƒrDtd ˆ¡ƒ‚‡ ‡fdd„}|S )a[  
    Decorator to register a pairwise function with :meth:`kl_divergence`.
    Usage::

        @register_kl(Normal, Normal)
        def kl_normal_normal(p, q):
            # insert implementation here

    Lookup returns the most specific (type,type) match ordered by subclass. If
    the match is ambiguous, a `RuntimeWarning` is raised. For example to
    resolve the ambiguous situation::

        @register_kl(BaseP, DerivedQ)
        def kl_version1(p, q): ...
        @register_kl(DerivedP, BaseQ)
        def kl_version2(p, q): ...

    you should register a third most-specific implementation, e.g.::

        register_kl(DerivedP, DerivedQ)(kl_version1)  # Break the tie.

    Args:
        type_p (type): A subclass of :class:`~torch.distributions.Distribution`.
        type_q (type): A subclass of :class:`~torch.distributions.Distribution`.
    z8Expected type_p to be a Distribution subclass but got {}z8Expected type_q to be a Distribution subclass but got {}c                    s   | t ˆ ˆf< t ¡  | S ©N)Ú_KL_REGISTRYr&   Úclear)Úfun©Útype_pÚtype_q© ú_/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/torch/distributions/kl.pyÚ	decoratorH   s    zregister_kl.<locals>.decorator)Ú
isinstanceÚtypeÚ
issubclassr   Ú	TypeErrorÚformat)r.   r/   r2   r0   r-   r1   r'   )   s    c                   @   s*   e Zd ZdgZdd„ Zdd„ Zdd„ ZdS )	Ú_MatchÚtypesc                 G   s
   || _ d S r)   ©r9   )Úselfr9   r0   r0   r1   Ú__init__T   s    z_Match.__init__c                 C   s   | j |j kS r)   r:   )r;   Úotherr0   r0   r1   Ú__eq__W   s    z_Match.__eq__c                 C   s8   t | j|jƒD ]$\}}t||ƒs& dS ||ur q4qdS )NFT)Úzipr9   r5   )r;   r=   ÚxÚyr0   r0   r1   Ú__le__Z   s    
z_Match.__le__N)Ú__name__Ú
__module__Ú__qualname__Ú	__slots__r<   r>   rB   r0   r0   r0   r1   r8   P   s   r8   c           	         s   ‡ ‡fdd„t D ƒ}|stS tdd„ |D ƒƒj\}}tdd„ |D ƒƒj\}}t ||f }t ||f }||urŒt d ˆ jˆj|j|j¡t¡ |S )zP
    Find the most specific approximate match, assuming single inheritance.
    c                    s,   g | ]$\}}t ˆ |ƒrt ˆ|ƒr||f‘qS r0   )r5   )Ú.0Zsuper_pZsuper_qr-   r0   r1   Ú
<listcomp>g   s   ÿz _dispatch_kl.<locals>.<listcomp>c                 s   s   | ]}t |Ž V  qd S r)   )r8   ©rG   Úmr0   r0   r1   Ú	<genexpr>n   ó    z_dispatch_kl.<locals>.<genexpr>c                 s   s   | ]}t t|ƒŽ V  qd S r)   )r8   ÚreversedrI   r0   r0   r1   rK   o   rL   z;Ambiguous kl_divergence({}, {}). Please register_kl({}, {}))	r*   ÚNotImplementedÚminr9   ÚwarningsÚwarnr7   rC   ÚRuntimeWarning)	r.   r/   ÚmatchesZleft_pZleft_qZright_qZright_pZleft_funZ	right_funr0   r-   r1   Ú_dispatch_klc   s    ÿþrT   c                 C   s   t  | t¡S )zI
    Helper function for obtaining infinite KL Divergence throughout
    )ÚtorchZ	full_liker   ©Ztensorr0   r0   r1   Ú_infinite_likey   s    rW   c                 C   s   | |   ¡  S )z2
    Utility function for calculating x log x
    )ÚlogrV   r0   r0   r1   Ú_x_log_x€   s    rY   c                 C   sD   |   d¡}|   d¡}|  d|| ¡ d¡ d¡}| | jdd… ¡S )zp
    Utility function for calculating the trace of XX^{T} with X having arbitrary trailing batch dimensions
    éÿÿÿÿéþÿÿÿé   N)ÚsizeZreshapeÚpowÚsumÚshape)ZbmatÚnrJ   Z
flat_tracer0   r0   r1   Ú_batch_trace_XXT‡   s    

rb   )ÚpÚqÚreturnc                 C   s|   zt t| ƒt|ƒf }W n8 tyP   tt| ƒt|ƒƒ}|t t| ƒt|ƒf< Y n0 |tu rrtd | jj|jj¡ƒ‚|| |ƒS )a"  
    Compute Kullback-Leibler divergence :math:`KL(p \| q)` between two distributions.

    .. math::

        KL(p \| q) = \int p(x) \log\frac {p(x)} {q(x)} \,dx

    Args:
        p (Distribution): A :class:`~torch.distributions.Distribution` object.
        q (Distribution): A :class:`~torch.distributions.Distribution` object.

    Returns:
        Tensor: A batch of KL divergences of shape `batch_shape`.

    Raises:
        NotImplementedError: If the distribution types have not been registered via
            :meth:`register_kl`.
    z8No KL(p || q) is implemented for p type {} and q type {})	r&   r4   ÚKeyErrorrT   rN   ÚNotImplementedErrorr7   Ú	__class__rC   )rc   rd   r,   r0   r0   r1   r(   ‘   s    ÿc                 C   s”   | j tjj |j ¡tjj | j ¡  }t||j dk< d|| j dk< d| j  tjj |j¡tjj | j¡  }t||j dk< d|| j dk< || S ©Nr   r   )ÚprobsrU   ÚnnZ
functionalZsoftplusÚlogitsr   ©rc   rd   Út1Út2r0   r0   r1   Ú_kl_bernoulli_bernoulli¶   s    **rp   c           	      C   s¦   | j | j }|j |j }|j  ¡ |j ¡  | ¡  }| j  ¡ | j ¡  | ¡  }| j |j  t | j ¡ }| j|j t | j¡ }|| t |¡ }|| | | | S r)   )Úconcentration1Úconcentration0ÚlgammarU   Údigamma)	rc   rd   Zsum_params_pZsum_params_qrn   ro   Út3Út4Út5r0   r0   r1   Ú_kl_beta_betaÁ   s    rx   c                 C   sh   | j |j k  ¡ rtdƒ‚| j | j| j|j  | j  ¡  |j  ¡   }| j |j k}t|| ƒ||< |S )NzKKL between Binomials where q.total_count > p.total_count is not implemented)Ztotal_countÚanyrg   rj   rl   Úlog1prW   )rc   rd   ÚklZinf_idxsr0   r0   r1   Ú_kl_binomial_binomialÍ   s    0r|   c                 C   sD   | j | j|j  }t||j dk |¡< d|| j dk |¡< | d¡S )Nr   rZ   )rj   rl   r   Z	expand_asr_   )rc   rd   Útr0   r0   r1   Ú_kl_categorical_categoricalÙ   s    r~   c                 C   sL   | j | j|j  }|  ¡ t | j ¡ }| ¡  t |j ¡ }|| | S r)   )Úmeanrl   Ú_cont_bern_log_normrU   rz   rj   ©rc   rd   rn   ro   ru   r0   r0   r1   Ú-_kl_continuous_bernoulli_continuous_bernoulliá   s    r‚   c                 C   s|   | j  d¡}|j  d¡}| ¡ | ¡  }| j  ¡ |j  ¡   d¡}| j |j  }| j  ¡ | ¡  d¡ }|| ||  d¡ S )NrZ   )Úconcentrationr_   rs   rt   Ú	unsqueeze)rc   rd   Zsum_p_concentrationZsum_q_concentrationrn   ro   ru   rv   r0   r0   r1   Ú_kl_dirichlet_dirichleté   s    r…   c                 C   s"   |j | j  }| ¡  }|| d S ©Nr   ©ÚraterX   )rc   rd   Z
rate_ratiorn   r0   r0   r1   Ú_kl_exponential_exponentialõ   s    
r‰   c                 C   s˜   t | ƒt |ƒkstdƒ‚dd„ | jD ƒ}|j}| j|Ž }tjj| ¡ |dd}|j|Ž | }t|||ƒD ]*\}}}	|| |	 }
|t	|
t
|jƒƒ8 }qh|S )Nz‡The cross KL-divergence between different exponential families cannot                             be computed using Bregman divergencesc                 S   s   g | ]}|  ¡  ¡ ‘qS r0   )ÚdetachZrequires_grad_)rG   Únpr0   r0   r1   rH     rL   z+_kl_expfamily_expfamily.<locals>.<listcomp>T)Zcreate_graph)r4   rg   Z_natural_paramsZ_log_normalizerrU   ZautogradZgradr_   r?   r$   ÚlenÚevent_shape)rc   rd   Z	p_nparamsZ	q_nparamsZ	lg_normalZ	gradientsÚresultZpnpZqnpÚgZtermr0   r0   r1   Ú_kl_expfamily_expfamilyü   s    
r   c                 C   sn   |j | j|j  ¡  }t |j ¡t | j ¡ }| j |j  t | j ¡ }|j| j | j | j  }|| | | S r)   )rƒ   rˆ   rX   rU   rs   rt   ©rc   rd   rn   ro   ru   rv   r0   r0   r1   Ú_kl_gamma_gamma  s
    r’   c                 C   sl   | j |j  }|j|j  }| j|j  }| ¡  | | }|t }t |d|  ¡  | ¡}|| | dt  S r†   )ÚscaleÚlocrX   Ú_euler_gammarU   Úexprs   )rc   rd   Zct1Zct2Zct3rn   ro   ru   r0   r0   r1   Ú_kl_gumbel_gumbel  s    r—   c                 C   s$   |   ¡  t |j ¡| j  |j S r)   )ÚentropyrU   rz   rj   rl   ©rc   rd   r0   r0   r1   Ú_kl_geometric_geometric   s    rš   c                 C   s   t | j|jƒS r)   )Ú_kl_normal_normalÚ	base_distr™   r0   r0   r1   Ú_kl_halfnormal_halfnormal%  s    r   c                 C   sV   | j |j  }| j|j  ¡ }| ¡  }||j  }|t | | j  ¡ }|| | d S r†   )r“   r”   ÚabsrX   rU   r–   )rc   rd   Úscale_ratioZloc_abs_diffrn   ro   ru   r0   r0   r1   Ú_kl_laplace_laplace*  s    

r    c                 C   sú   | j |j krtdƒ‚t|j|j|jƒt| j| j| jƒ }t|j|j|j| j |jƒ}|jj|j 	d¡ }t
jj|j|dd}| j|j  d¡}t| j|j ¡  	d¡ ƒ}t|| j ¡  	d¡ ƒ}t| | j¡ƒ}	|| | |	 }
d||
 | | j d   S )NzKL-divergence between two Low Rank Multivariate Normals with                          different event shapes cannot be computedr[   F©ÚupperrZ   ç      à?r   )r   Ú
ValueErrorr   Ú_unbroadcasted_cov_factorÚ_unbroadcasted_cov_diagÚ_capacitance_trilr   r”   ÚmTr„   rU   ÚlinalgÚsolve_triangularr_   rb   ÚrsqrtÚsqrtÚmatmul)rc   rd   Úterm1Úterm3Ú	qWt_qDinvÚAÚterm21Úterm22Zterm23Zterm24Úterm2r0   r0   r1   Ú7_kl_lowrankmultivariatenormal_lowrankmultivariatenormal5  s2    
ÿ
ÿþ

þ
ÿÿrµ   c           	      C   sÔ   | j |j krtdƒ‚t|j|j|jƒd| jjddd ¡  	d¡  }t
|j|j|j| j |jƒ}|jj|j d¡ }tjj|j|dd}t| j|j ¡  d¡ ƒ}t| | j¡ƒ}|| }d|| | | j d	   S )
NúKL-divergence between two (Low Rank) Multivariate Normals with                          different event shapes cannot be computedr\   r[   rZ   ©Zdim1Zdim2Fr¡   r£   r   )r   r¤   r   r¥   r¦   r§   Ú_unbroadcasted_scale_trilÚdiagonalrX   r_   r   r”   r¨   r„   rU   r©   rª   rb   r«   r­   )	rc   rd   r®   r¯   r°   r±   r²   r³   r´   r0   r0   r1   Ú0_kl_multivariatenormal_lowrankmultivariatenormalQ  s*    
ÿþ

þ
ÿÿrº   c                 C   s$  | j |j krtdƒ‚d|jjddd ¡  d¡ t| j| j| j	ƒ }t
|j|j| j ƒ}tj |jjd d… | jjd d… ¡}| j d }|j |||f ¡}| j ||| j d¡f ¡}t | j ¡ ¡ |||f ¡}ttjj||ddƒ}	ttjj||ddƒ}
|	|
 }d	|| | | j d   S )
Nr¶   r\   r[   rZ   r·   r   Fr¡   r£   )r   r¤   r¸   r¹   rX   r_   r   r¥   r¦   r§   r   r”   rU   Ú_CÚ_infer_sizer`   ÚexpandZ
cov_factorr]   Z
diag_embedr¬   rb   r©   rª   )rc   rd   r®   r¯   Úcombined_batch_shapera   Úq_scale_trilZp_cov_factorZ
p_cov_diagr²   r³   r´   r0   r0   r1   Ú0_kl_lowrankmultivariatenormal_multivariatenormalj  s.    
ÿÿÿ
ÿ
ÿrÀ   c           	      C   sÞ   | j |j krtdƒ‚|jjddd ¡  d¡| jjddd ¡  d¡ }tj |jj	d d… | jj	d d… ¡}| j d }|j 
|||f ¡}| j 
|||f ¡}ttjj||ddƒ}t|j|j| j ƒ}|d|| |   S )	NzvKL-divergence between two Multivariate Normals with                          different event shapes cannot be computedr[   rZ   r·   r   Fr¡   r£   )r   r¤   r¸   r¹   rX   r_   rU   r»   r¼   r`   r½   rb   r©   rª   r   r”   )	rc   rd   Z
half_term1r¾   ra   r¿   Zp_scale_trilr´   r¯   r0   r0   r1   Ú)_kl_multivariatenormal_multivariatenormal„  s    ÿÿ
rÁ   c                 C   sB   | j |j   d¡}| j|j |j   d¡}d|| d | ¡   S ©Nr\   r£   r   ©r“   r^   r”   rX   )rc   rd   Z	var_ratiorn   r0   r0   r1   r›   —  s    r›   c                 C   s   t | j|jƒS r)   )r~   Z_categoricalr™   r0   r0   r1   Ú'_kl_onehotcategorical_onehotcategoricalž  s    rÄ   c                 C   sX   | j |j  }|j| j }|j| ¡  }| ¡  }|| | d }t|| jj|jjk < |S r†   )r“   ÚalpharX   r   ÚsupportÚlower_bound)rc   rd   rŸ   Zalpha_ratiorn   ro   rŽ   r0   r0   r1   Ú_kl_pareto_pareto£  s    
rÈ   c                 C   s&   | j | j  ¡ |j  ¡   | j |j   S r)   r‡   r™   r0   r0   r1   Ú_kl_poisson_poisson¯  s    rÉ   c                 C   s.   | j |j krt‚| j|jkr t‚t| j|jƒS r)   )Z
transformsrg   r   r(   rœ   r™   r0   r0   r1   Ú_kl_transformed_transformed´  s
    rÊ   c                 C   s<   |j |j | j | j   ¡ }t||j| jk|j | j k B < |S r)   )ÚhighÚlowrX   r   ©rc   rd   rŽ   r0   r0   r1   Ú_kl_uniform_uniform½  s    rÎ   c                 C   s    |   ¡  | j|j ¡  |j  S r)   )r˜   rj   rˆ   rX   r™   r0   r0   r1   Ú_kl_bernoulli_poissonÅ  s    rÏ   c                 C   s,   |   ¡  | j|j  t |j ¡ | ¡  S r)   )r˜   r   rl   rU   rz   rj   r€   r™   r0   r0   r1   Ú_kl_beta_continuous_bernoulliÊ  s    rÐ   c                 C   s
   t | jƒS r)   )rW   rq   r™   r0   r0   r1   Ú_kl_beta_infinityÏ  s    rÑ   c                 C   s,   |   ¡  |j ¡  |j| j| j| j    S r)   )r˜   rˆ   rX   rq   rr   r™   r0   r0   r1   Ú_kl_beta_exponentialÔ  s    rÒ   c                 C   sp   |   ¡  }|j ¡ |j|j ¡   }|jd | j ¡ | j| j  ¡   }|j| j | j| j  }|| | | S r†   )r˜   rƒ   rs   rˆ   rX   rq   rt   rr   r‘   r0   r0   r1   Ú_kl_beta_gammaÙ  s
    
$rÓ   c           	      C   sš   | j | j | j  }|j d¡}|  ¡  }d|d tj  ¡  }|d|  | j | j d  | d¡ d }|j| }|j d¡d }|| || | |  S rÂ   )	rq   rr   r“   r^   r˜   ÚmathÚpirX   r”   )	rc   rd   ZE_betaÚ
var_normalrn   ro   ru   rv   rw   r0   r0   r1   Ú_kl_beta_normalä  s    
*
r×   c                 C   s>   |   ¡  |j|j  ¡  }t||j| jjk|j| jjk B < |S r)   )r˜   rË   rÌ   rX   r   rÆ   rÇ   Úupper_boundrÍ   r0   r0   r1   Ú_kl_beta_uniformð  s     rÙ   c                 C   s
   t | jƒS r)   )rW   rj   r™   r0   r0   r1   Ú!_kl_continuous_bernoulli_infinityù  s    rÚ   c                 C   s"   |   ¡  t |j¡ |j| j  S r)   )r˜   rU   rX   rˆ   r   r™   r0   r0   r1   Ú$_kl_continuous_bernoulli_exponentialþ  s    rÛ   c                 C   sz   |   ¡  }dt dtj ¡t |j|j ¡  t |j¡ }| jt | j	¡ d|j | j	  dt |j¡  }|| | S )Nr£   g       @)
r˜   rÔ   rX   rÕ   rU   Zsquarer”   r“   Zvariancer   r   r0   r0   r1   Ú_kl_continuous_bernoulli_normal  s    
22rÜ   c              	   C   sV   |   ¡  |j|j  ¡  }t t t |j| jj	¡t 
|j| jj¡¡t |¡t |¡S r)   )r˜   rË   rÌ   rX   rU   ÚwhereÚmaxÚgerÆ   rÇ   ÚlerØ   Ú	ones_liker   rÍ   r0   r0   r1   Ú _kl_continuous_bernoulli_uniform  s    ÿþrâ   c                 C   s
   t | jƒS r)   ©rW   rˆ   r™   r0   r0   r1   Ú_kl_exponential_infinity  s    rä   c                 C   sB   |j | j  }|j t |¡ }|| |j ¡  |jt  dt  S r†   )rˆ   rƒ   rU   rX   rs   r•   )rc   rd   Zratiorn   r0   r0   r1   Ú_kl_exponential_gamma  s    rå   c                 C   sR   | j |j }|j|j }| ¡ d }t |¡| |d  }| ¡ }|| | | S r†   )rˆ   r“   r”   rX   rU   r–   Ú
reciprocal)rc   rd   Úscale_rate_prodÚloc_scale_ratiorn   ro   ru   r0   r0   r1   Ú_kl_exponential_gumbel%  s    ré   c                 C   sp   |j  d¡}| j d¡}dt || d tj ¡ }| ¡ }|j| j }|j d¡d }|d || | |  S rÂ   )	r“   r^   rˆ   rU   rX   rÔ   rÕ   ræ   r”   )rc   rd   rÖ   Zrate_sqrrn   ro   ru   rv   r0   r0   r1   Ú_kl_exponential_normal1  s    rê   c                 C   s
   t | jƒS r)   )rW   rƒ   r™   r0   r0   r1   Ú_kl_gamma_infinity<  s    rë   c                 C   s&   |   ¡  |j ¡  |j| j | j  S r)   )r˜   rˆ   rX   rƒ   r™   r0   r0   r1   Ú_kl_gamma_exponentialD  s    rì   c                 C   s~   | j |j }|j|j }| jd | j ¡  | j ¡  | j }| ¡ | j|  }t |¡d| 	¡   
| j ¡ | }|| | S r†   )rˆ   r“   r”   rƒ   rt   rs   rX   rU   r–   ræ   r^   )rc   rd   Zbeta_scale_prodrè   rn   ro   ru   r0   r0   r1   Ú_kl_gamma_gumbelI  s    $$rí   c                 C   s¨   |j  d¡}| j d¡}dt || d tj ¡ | j | j ¡  }d| j d¡| j  | }|j	| j | j }d|j	 d¡ }|| jd | j 
¡   || | |  S rÂ   )r“   r^   rˆ   rU   rX   rÔ   rÕ   rƒ   rs   r”   rt   )rc   rd   rÖ   Zbeta_sqrrn   ro   ru   rv   r0   r0   r1   Ú_kl_gamma_normalU  s    ,rî   c                 C   s
   t | jƒS r)   ©rW   r”   r™   r0   r0   r1   Ú_kl_gumbel_infinity`  s    rð   c                 C   sx   | j |j  }|t dtj ¡  ¡ }tj| d  d¡d }| j| j t  |j |j   d¡d }| | | td  S )Nr\   r£   é   r   )r“   rÔ   r¬   rÕ   rX   r^   r”   r•   )rc   rd   Zparam_ratiorn   ro   ru   r0   r0   r1   Ú_kl_gumbel_normall  s
    &rò   c                 C   s
   t | jƒS r)   rï   r™   r0   r0   r1   Ú_kl_laplace_infinityu  s    ró   c                 C   s~   |j  d¡}| j  d¡| }dt d| tj ¡ }d| j d¡ }| j|j }d|j d¡ }| | || | |  d S rÂ   )r“   r^   rU   rX   rÔ   rÕ   r”   )rc   rd   rÖ   Zscale_sqr_var_ratiorn   ro   ru   rv   r0   r0   r1   Ú_kl_laplace_normal  s    rô   c                 C   s
   t | jƒS r)   rï   r™   r0   r0   r1   Ú_kl_normal_infinityŠ  s    rõ   c                 C   s|   | j |j }| j|j  d¡}|j |j }| ¡ d }|| }t | d|  | ¡}| | | ddt dtj ¡   S rÂ   )r”   r“   r^   rX   rU   r–   rÔ   rÕ   )rc   rd   Zmean_scale_ratioZvar_scale_sqr_ratiorè   rn   ro   ru   r0   r0   r1   Ú_kl_normal_gumbel”  s    rö   c                 C   sš   | j |j  }| j|j }|| j }t |¡}t dtj ¡| j t d| d¡ ¡ }|t 	t d¡| ¡ }| || |j  ddt dtj ¡   S )Nr\   g      à¿r£   r   )
r”   r“   rU   rX   rÔ   r¬   rÕ   r–   r^   Úerf)rc   rd   Zloc_diffrŸ   Zloc_diff_scale_ratiorn   ro   ru   r0   r0   r1   Ú_kl_normal_laplaceŸ  s    

*rø   c                 C   s
   t | jƒS r)   )rW   r“   r™   r0   r0   r1   Ú_kl_pareto_infinityª  s    rù   c                 C   sZ   | j |j }| j|  ¡ }| j ¡ }| j| | jd  }|| | d }t|| jdk< |S r†   )r“   rˆ   rÅ   rX   ræ   r   )rc   rd   rç   rn   ro   ru   rŽ   r0   r0   r1   Ú_kl_pareto_exponential±  s    
rú   c                 C   sŒ   | j  ¡ | j ¡  }| j ¡ | }|j ¡ |j|j ¡   }d|j | }|j| j | j  | jd  }|| | | d }t|| jdk< |S r†   )r“   rX   rÅ   ræ   rƒ   rs   rˆ   r   ©rc   rd   Úcommon_termrn   ro   ru   rv   rŽ   r0   r0   r1   Ú_kl_pareto_gamma¼  s    rý   c           	      C   sª   d|j  d¡ }| j | jd  }t dtj ¡|j  | j | j   ¡ }| j ¡ }| j| d¡ | jd  }| j| |j  d¡}|| || |  d }t	|| jdk< |S )Nr\   r   )
r“   r^   rÅ   rÔ   r¬   rÕ   rX   ræ   r”   r   )	rc   rd   rÖ   rü   rn   ro   ru   rv   rŽ   r0   r0   r1   Ú_kl_pareto_normalÊ  s    &
rþ   c                 C   s
   t | jƒS r)   rã   r™   r0   r0   r1   Ú_kl_poisson_infinity×  s    rÿ   c                 C   sÂ   | j | j }t |¡}|jd t| j ƒt| jƒ |  | }|jd td| j  ƒtd| j ƒ |  | }|j ¡ |j ¡  |j|j  ¡  }|| | | }t|| j |j	j
k| j|j	jk B < |S r†   )rË   rÌ   rU   rX   rq   rY   rr   rs   r   rÆ   rØ   rÇ   rû   r0   r0   r1   Ú_kl_uniform_betaÝ  s    
&.$ r   c              	   C   sh   |   ¡  | j|j  t |j ¡ | ¡  }t t t 	| j
|jj¡t | j|jj¡¡t |¡t |¡S r)   )r˜   r   rl   rU   rz   rj   r€   rÝ   rÞ   rß   rË   rÆ   rØ   rà   rÌ   rÇ   rá   r   rÍ   r0   r0   r1   Ú _kl_uniform_continuous_bernoullié  s    ,ÿþr  c                 C   sB   |j | j| j  d | j| j |j   ¡  }t|| j|jjk < |S )Nr\   )rˆ   rË   rÌ   rX   r   rÆ   rÇ   rÍ   r0   r0   r1   Ú_kl_uniform_exponetialñ  s    ,r  c                 C   s’   | j | j }| ¡ }|j ¡ |j|j ¡   }d|j t| j ƒt| jƒ |  | }|j| j | j  d }| | | | }t|| j|jj	k < |S )Nr   r\   )
rË   rÌ   rX   rƒ   rs   rˆ   rY   r   rÆ   rÇ   rû   r0   r0   r1   Ú_kl_uniform_gammaø  s    &r  c                 C   sn   |j | j| j  }| j|j |j  }| j|j |j  }| ¡ d||   }|t | ¡t | ¡  }|| S )Nr£   )r“   rË   rÌ   r”   rX   rU   r–   )rc   rd   rü   Zhigh_loc_diffZlow_loc_diffrn   ro   r0   r0   r1   Ú_kl_uniform_gumbel  s    r  c                 C   st   | j | j }t tjd ¡|j |  ¡ }| d¡d }| j | j d|j  d  d¡}|d||  |j d¡  S )Nr\   é   r£   )	rË   rÌ   rÔ   r¬   rÕ   r“   rX   r^   r”   )rc   rd   rü   rn   ro   ru   r0   r0   r1   Ú_kl_uniform_normal  s
     r  c                 C   sl   | j | j }|j|j |j¡ |  ¡ }t| j ƒt| jƒ | | }||jd  | }t|| j|jj	k < |S r†   )
rË   rÌ   rÅ   r“   r^   rX   rY   r   rÆ   rÇ   )rc   rd   Zsupport_uniformrn   ro   rŽ   r0   r0   r1   Ú_kl_uniform_pareto  s    r  c                 C   s*   | j |j krt‚t| j|jƒ}t|| j ƒS r)   )Zreinterpreted_batch_ndimsrg   r(   rœ   r$   rÍ   r0   r0   r1   Ú_kl_independent_independent#  s    r  c                 C   sD   | j |j   d¡| j|j  d¡  ¡ }d| j  |j   ¡ }|| S )Nr\   é   rÃ   rm   r0   r0   r1   Ú_kl_cauchy_cauchy+  s    (r
  c                  C   sX   dg} t tdd„ dD ]\}}|  d |j|j¡¡ qd | ¡}tjrTt j|7  _dS )zHAppends a list of implemented KL functions to the doc for kl_divergence.zLKL divergence is currently implemented for the following distribution pairs:c                 S   s   | d j | d j fS ri   )rC   )Zp_qr0   r0   r1   Ú<lambda>6  rL   z_add_kl_info.<locals>.<lambda>)ÚkeyzG* :class:`~torch.distributions.{}` and :class:`~torch.distributions.{}`z
	N)Úsortedr*   Úappendr7   rC   Újoinr(   Ú__doc__)Úrowsrc   rd   Zkl_infor0   r0   r1   Ú_add_kl_info2  s    ÿÿ
r  )ŽrÔ   rP   Ú	functoolsr   Útypingr   r   r   r   rU   r   Z	bernoullir	   Úbetar
   Zbinomialr   Zcategoricalr   Zcauchyr   Zcontinuous_bernoullir   Z	dirichletr   Údistributionr   Zexponentialr   Z
exp_familyr   Úgammar   Z	geometricr   Zgumbelr   Zhalf_normalr   Zindependentr   Zlaplacer   Zlowrank_multivariate_normalr   r   r   Zmultivariate_normalr   r   Únormalr   Zone_hot_categoricalr   Zparetor    Zpoissonr!   Ztransformed_distributionr"   Úuniformr#   Úutilsr$   r%   r•   r*   r&   Ú__annotations__Ú__all__r'   r8   rT   rW   rY   rb   ZTensorr(   rp   rx   r|   r~   r‚   r…   r‰   r   r’   r—   rš   r   r    rµ   rº   rÀ   rÁ   r›   rÄ   rÈ   rÉ   rÊ   rÎ   rÏ   rÐ   rÑ   rÒ   rÓ   r×   rÙ   rÚ   rÛ   rÜ   râ   rä   rå   ré   rê   rë   rì   rí   rî   rð   rò   ró   rô   rõ   rö   rø   rù   rú   rý   rþ   rÿ   r   r  r  r  r  r  r  r  r
  r  r0   r0   r0   r1   Ú<module>   s€  
'
%































































	

