MCPcopy Create free account
hub / github.com/VisionXLab/OF-Diff / normal_kl

Function normal_kl

ldm/modules/distributions/distributions.py:65–92  ·  view source on GitHub ↗

source: https://github.com/openai/guided-diffusion/blob/27c20a8fab9cb472df5d6bdd6c8d11c8f430b924/guided_diffusion/losses.py#L12 Compute the KL divergence between two gaussians. Shapes are automatically broadcasted, so batches can be compared to scalars, among other use cases.

(mean1, logvar1, mean2, logvar2)

Source from the content-addressed store, hash-verified

63
64
65def normal_kl(mean1, logvar1, mean2, logvar2):
66 """
67 source: https://github.com/openai/guided-diffusion/blob/27c20a8fab9cb472df5d6bdd6c8d11c8f430b924/guided_diffusion/losses.py#L12
68 Compute the KL divergence between two gaussians.
69 Shapes are automatically broadcasted, so batches can be compared to
70 scalars, among other use cases.
71 """
72 tensor = None
73 for obj in (mean1, logvar1, mean2, logvar2):
74 if isinstance(obj, torch.Tensor):
75 tensor = obj
76 break
77 assert tensor is not None, "at least one argument must be a Tensor"
78
79 # Force variances to be Tensors. Broadcasting helps convert scalars to
80 # Tensors, but it does not work for torch.exp().
81 logvar1, logvar2 = [
82 x if isinstance(x, torch.Tensor) else torch.tensor(x).to(tensor)
83 for x in (logvar1, logvar2)
84 ]
85
86 return 0.5 * (
87 -1.0
88 + logvar2
89 - logvar1
90 + torch.exp(logvar1 - logvar2)
91 + ((mean1 - mean2) ** 2) * torch.exp(-logvar2)
92 )

Callers 1

_prior_bpdMethod · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected