Team Ai
Apppublic

wkea/blockdiffusion-api

sourceHugging Faceupdated 9mo agoView on Hugging Face
0likes
gaussian_diffusion.py254 linesDownload Raw Back to diffusion
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