wkea/blockdiffusion-api
0
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 