-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtest_ddim.py
More file actions
126 lines (99 loc) · 4.61 KB
/
Copy pathtest_ddim.py
File metadata and controls
126 lines (99 loc) · 4.61 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
"""
Inference Script Version Apr 17th 2023
"""
from dataclasses import dataclass
from tools import *
from unet import *
#from DDIM import *
from DDIM_new import *
from IPython.display import display, HTML
from torch.utils.data import DataLoader
import random
from tqdm import tqdm
seed = 3407
random.seed(seed)
torch.manual_seed(seed)
os.environ['PYTHONHASHSEED'] = str(seed)
np.random.seed(seed)
torch.cuda.manual_seed(seed)
torch.backends.cudnn.deterministic = True
torch.backends.cudnn.benchmark = False
@dataclass
class BaseConfig:
DEVICE = get_default_device()
DATASET = "MNIST" # "Cifar-10", "Cifar-100", "Flowers"
# Path to log inference images and save checkpoints
root = "./Logs_Checkpoints"
os.makedirs(root, exist_ok=True)
# Current log and checkpoint directory.
log_folder = 'version_4' # specific a folder name to load
checkpoint_name = "ddim_1.tar"
@dataclass
class TrainingConfig:
TIMESTEPS = 1000 # Define number of diffusion timesteps
IMG_SHAPE = (1, 32, 32) if BaseConfig.DATASET == "MNIST" else (3, 32, 32)
NUM_EPOCHS = 25
BATCH_SIZE = 16
LR = 2e-4
NUM_WORKERS = 4 if str(BaseConfig.DEVICE) != "cpu" else 0
@dataclass
class ModelConfig: # setting up attention unet
BASE_CH = 64
BASE_CH_MULT = (1, 2, 4, 8)
APPLY_ATTENTION = (False, False, True, False)
DROPOUT_RATE = 0.1
TIME_EMB_MULT = 2
sd = Diffusion_setting(num_diffusion_timesteps=TrainingConfig.TIMESTEPS,
img_shape=TrainingConfig.IMG_SHAPE, device=BaseConfig.DEVICE)
generate_video = False
# test
log_dir, checkpoint_dir = setup_log_directory(config=BaseConfig(), inference=True)
model = UNet(
input_channels=TrainingConfig.IMG_SHAPE[0],
output_channels=TrainingConfig.IMG_SHAPE[0],
base_channels=ModelConfig.BASE_CH,
base_channels_multiples=ModelConfig.BASE_CH_MULT,
apply_attention=ModelConfig.APPLY_ATTENTION,
dropout_rate=ModelConfig.DROPOUT_RATE,
time_multiple=ModelConfig.TIME_EMB_MULT,
)
model.load_state_dict(torch.load(os.path.join(checkpoint_dir, BaseConfig.checkpoint_name), map_location='cpu')["model"], False)
model.to(BaseConfig.DEVICE)
# 获取数据集
dataset = get_dataset(dataset_name=BaseConfig.DATASET)
dataloader = DataLoader(dataset, batch_size=16, shuffle=False)
device_dataloader = DeviceDataLoader(dataloader, BaseConfig.DEVICE)
# 获取一批数据
X_0_batch, _ = next(iter(device_dataloader))
# 展示原图
original_images = make_a_grid_based_PIL_npy(X_0_batch, nrow=4)
original_images.save(os.path.join(log_dir, "original_images.png"))
display(original_images)
# 扩散过程
timesteps = torch.full((X_0_batch.shape[0],), TrainingConfig.TIMESTEPS-1, device=BaseConfig.DEVICE, dtype=torch.long)
X_t_batch, _ = forward_diffusion(sd, X_0_batch, timesteps)
# 展示扩散后的图
noisy_images = make_a_grid_based_PIL_npy(X_t_batch, nrow=4)
noisy_images.save(os.path.join(log_dir, "noisy_images.png"))
display(noisy_images)
os.makedirs(log_dir, exist_ok=True)
ext = ".mp4" if generate_video else ".png"
filename = f"{datetime.now().strftime('%Y%m%d-%H%M%S')}{ext}"
save_path = os.path.join(log_dir, filename)
# 还原过程
# denoising_reverse_diffusion(model, sd, x_T= X_t_batch, img_shape=TrainingConfig.IMG_SHAPE, num_images=X_0_batch.shape[0], nrow=4,
# save_path=save_path, generate_video=generate_video, device=BaseConfig.DEVICE, eta=0, tau=5)
# probablity_flow(model, sd, x_T= X_t_batch, img_shape=TrainingConfig.IMG_SHAPE, num_images=X_0_batch.shape[0], nrow=4,
# save_path=save_path, generate_video=generate_video, device=BaseConfig.DEVICE, eta=0, tau=5)
x_diffusion = probability_flow(model, sd, img_shape=TrainingConfig.IMG_SHAPE, num_images=X_0_batch.shape[0], nrow=4,
save_path=save_path, generate_video=generate_video, device=BaseConfig.DEVICE, eta=1, tau=1,reverse=True, x_t=X_0_batch)
#正向,得到的是噪声图
filename = f"{datetime.now().strftime('%Y%m%d-%H%M%S')}{ext}"
save_path = os.path.join(log_dir, filename)
probability_flow(model, sd, img_shape=TrainingConfig.IMG_SHAPE, num_images=X_0_batch.shape[0], nrow=4,
save_path=save_path, generate_video=generate_video, device=BaseConfig.DEVICE, eta=1, tau=1,reverse=False, x_t=x_diffusion)
filename = f"{datetime.now().strftime('%Y%m%d-%H%M%S')}{ext}"
save_path = os.path.join(log_dir, filename)
probability_flow(model, sd, img_shape=TrainingConfig.IMG_SHAPE, num_images=X_0_batch.shape[0], nrow=4,
save_path=save_path, generate_video=generate_video, device=BaseConfig.DEVICE, eta=1, tau=1,reverse=False, x_t=X_t_batch)
# 我怀疑这个函数有问题,我明天要看看这个函数的输出