Team Ai
Apppublic

wkea/blockdiffusion-api

sourceHugging Faceupdated 9mo agoView on Hugging Face
0likes
schedules.py44 linesDownload Raw Back to diffusion
1from __future__ import annotations
2
3import math
4from typing import Literal
5
6import torch
7
8
9def linear_beta_schedule(timesteps: int, beta_start: float = 1e-4, beta_end: float = 2e-2) -> torch.Tensor:
10    """
11    经典线性 beta 调度(DDPM)。
12    """
13    return torch.linspace(beta_start, beta_end, timesteps, dtype=torch.float32)
14
15
16def cosine_beta_schedule(timesteps: int, s: float = 0.008) -> torch.Tensor:
17    """
18    Cosine 调度(Nichol & Dhariwal, 2021)。
19    参考实现:alpha_bar(t) = cos^2((t+s)/(1+s)*pi/2)
20    """
21    steps = timesteps + 1
22    x = torch.linspace(0, timesteps, steps, dtype=torch.float64)
23    t = x / timesteps
24    alpha_bar = torch.cos((t + s) / (1 + s) * math.pi / 2) ** 2
25    alpha_bar = alpha_bar / alpha_bar[0]
26    betas = 1 - (alpha_bar[1:] / alpha_bar[:-1])
27    return betas.clamp(1e-8, 0.999).float()
28
29
30def get_beta_schedule(
31    schedule: Literal["linear", "cosine"],
32    timesteps: int,
33) -> torch.Tensor:
34    """
35    统一入口:返回长度为 timesteps 的 betas。
36    """
37    if schedule == "linear":
38        return linear_beta_schedule(timesteps)
39    if schedule == "cosine":
40        return cosine_beta_schedule(timesteps)
41    raise ValueError(f"未知 schedule: {schedule}")
42
43
44