wkea/blockdiffusion-api
0
1from __future__ import annotations2 3from dataclasses import dataclass4from typing import Any, Dict, Optional, Tuple5 6import torch7from torch import nn8 9 10def _extract(a: torch.Tensor, t: torch.Tensor, x_shape: Tuple[int, ...]) -> torch.Tensor:11 """12 从 [T] 的一维张量 a 中按 batch 的 t 抽取,并 reshape 到 x 的维度用于广播。13 """14 b = t.shape[0]15 out = a.gather(0, t)16 return out.reshape(b, *((1,) * (len(x_shape) - 1)))17 18 19@dataclass(frozen=True)20class DiffusionConfig:21 """22 扩散配置。23 """24 25 timesteps: int = 100026 beta_schedule: str = "cosine" # "linear" 或 "cosine"27 28 29class GaussianDiffusion:30 """31 标准 Gaussian Diffusion:32 - 训练:epsilon-prediction,loss 为 MSE33 - 采样:支持 DDPM 与(可选)DDIM34 训练循环/保存等不在这里做,以保证训练与采样逻辑解耦。35 """36 37 def __init__(self, betas: torch.Tensor):38 if betas.ndim != 1:39 raise ValueError("betas 必须是一维张量")40 self.device = betas.device41 self.betas = betas.float()42 self.num_timesteps = int(self.betas.shape[0])43 44 alphas = 1.0 - self.betas45 alphas_cumprod = torch.cumprod(alphas, dim=0)46 alphas_cumprod_prev = torch.cat([torch.tensor([1.0], device=betas.device), alphas_cumprod[:-1]], dim=0)47 48 self.alphas = alphas49 self.alphas_cumprod = alphas_cumprod50 self.alphas_cumprod_prev = alphas_cumprod_prev51 52 self.sqrt_alphas_cumprod = torch.sqrt(alphas_cumprod)53 self.sqrt_one_minus_alphas_cumprod = torch.sqrt(1.0 - alphas_cumprod)54 self.sqrt_recip_alphas_cumprod = torch.sqrt(1.0 / alphas_cumprod)55 self.sqrt_recipm1_alphas_cumprod = torch.sqrt(1.0 / alphas_cumprod - 1.0)56 57 # q(x_{t-1}|x_t,x_0) 的 posterior58 self.posterior_variance = (59 self.betas * (1.0 - alphas_cumprod_prev) / (1.0 - alphas_cumprod)60 ).clamp(min=1e-20)61 self.posterior_log_variance_clipped = torch.log(self.posterior_variance)62 self.posterior_mean_coef1 = self.betas * torch.sqrt(alphas_cumprod_prev) / (1.0 - alphas_cumprod)63 self.posterior_mean_coef2 = (1.0 - alphas_cumprod_prev) * torch.sqrt(alphas) / (1.0 - alphas_cumprod)64 65 def q_sample(self, x_start: torch.Tensor, t: torch.Tensor, noise: Optional[torch.Tensor] = None) -> torch.Tensor:66 """67 前向加噪:q(x_t | x_0)68 """69 if noise is None:70 noise = torch.randn_like(x_start)71 return _extract(self.sqrt_alphas_cumprod, t, x_start.shape) * x_start + _extract(72 self.sqrt_one_minus_alphas_cumprod, t, x_start.shape73 ) * noise74 75 def predict_x0_from_eps(self, x_t: torch.Tensor, t: torch.Tensor, eps: torch.Tensor) -> torch.Tensor:76 return _extract(self.sqrt_recip_alphas_cumprod, t, x_t.shape) * x_t - _extract(77 self.sqrt_recipm1_alphas_cumprod, t, x_t.shape78 ) * eps79 80 def q_posterior(self, x_start: torch.Tensor, x_t: torch.Tensor, t: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:81 """82 返回 q(x_{t-1}|x_t,x_0) 的 posterior mean 与 variance。83 """84 posterior_mean = _extract(self.posterior_mean_coef1, t, x_t.shape) * x_start + _extract(85 self.posterior_mean_coef2, t, x_t.shape86 ) * x_t87 posterior_variance = _extract(self.posterior_variance, t, x_t.shape)88 return posterior_mean, posterior_variance89 90 def training_losses(91 self,92 model: nn.Module,93 x_start: torch.Tensor,94 t: torch.Tensor,95 model_kwargs: Optional[Dict[str, Any]] = None,96 cond_drop_prob: float = 0.0,97 snr_gamma: float = 0.0,98 ) -> Dict[str, torch.Tensor]:99 """100 训练损失:MSE(eps_pred, eps)。101 可选:SNR 加权(snr_gamma>0),通常能提升采样质量/稳定性。102 """103 model_kwargs = dict(model_kwargs or {})104 noise = torch.randn_like(x_start)105 x_t = self.q_sample(x_start, t, noise=noise)106 107 # Classifier-Free Guidance 的训练方式:以一定概率丢弃条件(mask 全 0 => 等价无条件)108 if cond_drop_prob and cond_drop_prob > 0 and ("mask" in model_kwargs):109 mask = model_kwargs.get("mask", None)110 if mask is not None:111 # mask: (B,L);drop: (B,1) 广播到 (B,L)112 drop = torch.rand((x_start.shape[0], 1), device=x_start.device) < float(cond_drop_prob)113 if mask.dtype == torch.bool:114 model_kwargs["mask"] = mask & (~drop)115 else:116 drop_f = drop.to(dtype=mask.dtype)117 model_kwargs["mask"] = mask * (1 - drop_f)118 119 eps_pred = model(x_t, t, **model_kwargs)120 mse = (eps_pred - noise) ** 2 # (B,C,H,W)121 loss_per = mse.mean(dim=(1, 2, 3)) # (B,)122 123 if snr_gamma and snr_gamma > 0:124 # snr = alpha_bar / (1-alpha_bar)125 alpha_bar = _extract(self.alphas_cumprod, t, x_start.shape).reshape(x_start.shape[0])126 snr = alpha_bar / (1.0 - alpha_bar).clamp(min=1e-12)127 # 经典做法:w = min(snr, gamma) / snr128 w = torch.minimum(snr, torch.full_like(snr, float(snr_gamma))) / snr.clamp(min=1e-12)129 loss_per = loss_per * w130 131 loss = loss_per.mean()132 return {"loss": loss}133 134 @torch.no_grad()135 def p_mean_variance(136 self,137 model: nn.Module,138 x_t: torch.Tensor,139 t: torch.Tensor,140 model_kwargs: Optional[Dict[str, Any]] = None,141 ) -> Dict[str, torch.Tensor]:142 """143 DDPM 反向一步:根据 eps_pred 推出 x0_pred,再得到 posterior。144 """145 model_kwargs = dict(model_kwargs or {})146 eps_pred = model(x_t, t, **model_kwargs)147 x0_pred = self.predict_x0_from_eps(x_t, t, eps_pred).clamp(-1.0, 1.0)148 model_mean, model_var = self.q_posterior(x0_pred, x_t, t)149 return {"mean": model_mean, "variance": model_var, "pred_x0": x0_pred, "pred_eps": eps_pred}150 151 @torch.no_grad()152 def p_sample(153 self,154 model: nn.Module,155 x_t: torch.Tensor,156 t: torch.Tensor,157 model_kwargs: Optional[Dict[str, Any]] = None,158 ) -> torch.Tensor:159 """160 采样一步:x_t -> x_{t-1}(DDPM)。161 """162 out = self.p_mean_variance(model, x_t, t, model_kwargs=model_kwargs)163 mean, var = out["mean"], out["variance"]164 noise = torch.randn_like(x_t)165 # t==0 时不再加噪声166 nonzero_mask = (t != 0).float().reshape(x_t.shape[0], *((1,) * (x_t.ndim - 1)))167 return mean + nonzero_mask * torch.sqrt(var) * noise168 169 @torch.no_grad()170 def p_sample_loop(171 self,172 model: nn.Module,173 shape: Tuple[int, int, int, int],174 device: torch.device,175 model_kwargs: Optional[Dict[str, Any]] = None,176 progress: bool = True,177 ) -> torch.Tensor:178 """179 从纯噪声开始逐步采样到 x_0(DDPM)。180 """181 img = torch.randn(shape, device=device)182 timesteps = range(self.num_timesteps - 1, -1, -1)183 if progress:184 try:185 from tqdm import tqdm # 延迟导入,避免强依赖186 except Exception:187 tqdm = None188 if tqdm is not None:189 timesteps = tqdm(timesteps, desc="DDPM Sampling", total=self.num_timesteps)190 191 for i in timesteps:192 t = torch.full((shape[0],), i, device=device, dtype=torch.long)193 img = self.p_sample(model, img, t, model_kwargs=model_kwargs)194 return img195 196 @torch.no_grad()197 def ddim_sample_loop(198 self,199 model: nn.Module,200 shape: Tuple[int, int, int, int],201 device: torch.device,202 steps: int = 50,203 eta: float = 0.0,204 model_kwargs: Optional[Dict[str, Any]] = None,205 progress: bool = True,206 ) -> torch.Tensor:207 """208 DDIM 采样(更少步数、更快)。eta=0 为确定性采样。209 """210 if steps <= 1:211 raise ValueError("steps 必须 > 1")212 img = torch.randn(shape, device=device)213 214 # 选择等间隔的时间序列(从 T-1 到 0)215 # 注意:这里用 Python int 列表,避免在 CUDA Tensor 上做 int() 转换导致报错。216 times = torch.linspace(self.num_timesteps - 1, 0, steps, dtype=torch.float32).long().tolist()217 218 if progress:219 try:220 from tqdm import tqdm221 except Exception:222 tqdm = None223 if tqdm is not None:224 times = tqdm(list(times), desc="DDIM Sampling", total=steps)225 226 for idx, t_i in enumerate(times):227 t = torch.full((shape[0],), int(t_i), device=device, dtype=torch.long)228 out = self.p_mean_variance(model, img, t, model_kwargs=model_kwargs)229 eps = out["pred_eps"]230 x0 = out["pred_x0"]231 232 if idx == len(times) - 1:233 img = x0234 break235 236 t_next = torch.full((shape[0],), int(times[idx + 1]), device=device, dtype=torch.long)237 238 a_t = _extract(self.alphas_cumprod, t, img.shape)239 a_next = _extract(self.alphas_cumprod, t_next, img.shape)240 241 # DDIM 更新公式242 sigma = (243 eta244 * torch.sqrt((1.0 - a_next) / (1.0 - a_t))245 * torch.sqrt(1.0 - a_t / a_next)246 )247 noise = torch.randn_like(img)248 dir_xt = torch.sqrt(1.0 - a_next - sigma**2) * eps249 img = torch.sqrt(a_next) * x0 + dir_xt + sigma * noise250 251 return img252 253 254 