mirror of
https://github.com/wassname/denoising-diffusion-pytorch.git
synced 2026-09-12 12:22:11 +08:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
12079cadee |
@@ -440,7 +440,7 @@ class GaussianDiffusion(nn.Module):
|
|||||||
|
|
||||||
self.objective = objective
|
self.objective = objective
|
||||||
|
|
||||||
assert objective in {'pred_noise', 'pred_x0'}, 'objective must be either pred_noise (predict noise) or pred_x0 (predict image start)'
|
assert objective in {'pred_noise', 'pred_x0', 'pred_v'}, 'objective must be either pred_noise (predict noise) or pred_x0 (predict image start) or pred_v (predict v [v-parameterization as defined in appendix D of progressive distillation paper, used in imagen-video successfully])'
|
||||||
|
|
||||||
if beta_schedule == 'linear':
|
if beta_schedule == 'linear':
|
||||||
betas = linear_beta_schedule(timesteps)
|
betas = linear_beta_schedule(timesteps)
|
||||||
@@ -511,6 +511,18 @@ class GaussianDiffusion(nn.Module):
|
|||||||
extract(self.sqrt_recipm1_alphas_cumprod, t, x_t.shape)
|
extract(self.sqrt_recipm1_alphas_cumprod, t, x_t.shape)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def predict_v(self, x_start, t, noise):
|
||||||
|
return (
|
||||||
|
extract(self.sqrt_alphas_cumprod, t, x_start.shape) * noise -
|
||||||
|
extract(self.sqrt_one_minus_alphas_cumprod, t, x_start.shape) * x_start
|
||||||
|
)
|
||||||
|
|
||||||
|
def predict_start_from_v(self, x_t, t, v):
|
||||||
|
return (
|
||||||
|
extract(self.sqrt_alphas_cumprod, t, x_t.shape) * x_t -
|
||||||
|
extract(self.sqrt_one_minus_alphas_cumprod, t, x_t.shape) * v
|
||||||
|
)
|
||||||
|
|
||||||
def q_posterior(self, x_start, x_t, t):
|
def q_posterior(self, x_start, x_t, t):
|
||||||
posterior_mean = (
|
posterior_mean = (
|
||||||
extract(self.posterior_mean_coef1, t, x_t.shape) * x_start +
|
extract(self.posterior_mean_coef1, t, x_t.shape) * x_start +
|
||||||
@@ -534,6 +546,12 @@ class GaussianDiffusion(nn.Module):
|
|||||||
x_start = maybe_clip(x_start)
|
x_start = maybe_clip(x_start)
|
||||||
pred_noise = self.predict_noise_from_start(x, t, x_start)
|
pred_noise = self.predict_noise_from_start(x, t, x_start)
|
||||||
|
|
||||||
|
elif self.objective == 'pred_v':
|
||||||
|
v = model_output
|
||||||
|
x_start = self.predict_start_from_v(x, t, v)
|
||||||
|
x_start = maybe_clip(x_start)
|
||||||
|
pred_noise = self.predict_noise_from_start(x, t, x_start)
|
||||||
|
|
||||||
return ModelPrediction(pred_noise, x_start)
|
return ModelPrediction(pred_noise, x_start)
|
||||||
|
|
||||||
def p_mean_variance(self, x, t, x_self_cond = None, clip_denoised = True):
|
def p_mean_variance(self, x, t, x_self_cond = None, clip_denoised = True):
|
||||||
@@ -671,6 +689,9 @@ class GaussianDiffusion(nn.Module):
|
|||||||
target = noise
|
target = noise
|
||||||
elif self.objective == 'pred_x0':
|
elif self.objective == 'pred_x0':
|
||||||
target = x_start
|
target = x_start
|
||||||
|
elif self.objective == 'pred_v':
|
||||||
|
v = self.predict_v(x_start, t, noise)
|
||||||
|
target = v
|
||||||
else:
|
else:
|
||||||
raise ValueError(f'unknown objective {self.objective}')
|
raise ValueError(f'unknown objective {self.objective}')
|
||||||
|
|
||||||
|
|||||||
@@ -3,7 +3,7 @@ from setuptools import setup, find_packages
|
|||||||
setup(
|
setup(
|
||||||
name = 'denoising-diffusion-pytorch',
|
name = 'denoising-diffusion-pytorch',
|
||||||
packages = find_packages(),
|
packages = find_packages(),
|
||||||
version = '0.29.1',
|
version = '0.30.0',
|
||||||
license='MIT',
|
license='MIT',
|
||||||
description = 'Denoising Diffusion Probabilistic Models - Pytorch',
|
description = 'Denoising Diffusion Probabilistic Models - Pytorch',
|
||||||
author = 'Phil Wang',
|
author = 'Phil Wang',
|
||||||
|
|||||||
Reference in New Issue
Block a user