-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathgenerate.py
More file actions
66 lines (50 loc) · 2.25 KB
/
Copy pathgenerate.py
File metadata and controls
66 lines (50 loc) · 2.25 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
import torch
from config import DEVICE, EMBEDDING_SIZE, INNER_DIM, TIMESTEPS, MEAN, STD
from models import VAE, Diffuser
from PIL import Image
import numpy as np
@torch.no_grad()
def sample_minecraft(model, vae, num_steps=50, eta=0.0):
model.eval()
vae.eval()
# 1. Start with random noise
xt = torch.randn(1, 880, EMBEDDING_SIZE).to(DEVICE)
# 2. Recreate Cosine schedule
s = 0.008
t_seq = torch.linspace(0, num_steps, num_steps + 1).to(DEVICE)
alphas_cumprod = torch.cos(((t_seq / num_steps) + s) / (1 + s) * torch.pi / 2)**2
alphas_cumprod = alphas_cumprod / alphas_cumprod[0]
# 3. Denoising loop (DDIM)
for i in reversed(range(num_steps)):
t_batch = torch.full((1,), i, device=DEVICE, dtype=torch.long)
eps_pred = model(xt, t_batch)
alpha_t = alphas_cumprod[i+1]
alpha_tm1 = alphas_cumprod[i]
# Predict x0
x0_pred = (xt - torch.sqrt(1 - alpha_t) * eps_pred) / torch.sqrt(alpha_t)
x0_pred = x0_pred.clamp(-2.0, 2.0)
# Direction to xt-1
sigma_t = eta * torch.sqrt((1 - alpha_tm1) / (1 - alpha_t)) * torch.sqrt(1 - alpha_t / alpha_tm1)
direction_xt = torch.sqrt(1 - alpha_tm1 - sigma_t**2) * eps_pred
xt = torch.sqrt(alpha_tm1) * x0_pred + direction_xt
if i > 0 and eta > 0:
noise = torch.randn_like(xt)
xt = xt + sigma_t * noise
# 4. Decode latent to image
img_out = vae.decode(xt)
# Denormalize
mean = torch.tensor(MEAN).view(3, 1, 1).to(DEVICE)
std = torch.tensor(STD).view(3, 1, 1).to(DEVICE)
img_out = (img_out[0] * std + mean).clamp(0, 1)
return img_out.cpu().numpy().transpose(1, 2, 0)
if __name__ == "__main__":
# Load models
vae = VAE(dim=EMBEDDING_SIZE).to(DEVICE)
vae.load_state_dict(torch.load('vae_4.pt', map_location=DEVICE))
diffuser = Diffuser(dim=EMBEDDING_SIZE, inner_dim=INNER_DIM).to(DEVICE)
diffuser.load_state_dict(torch.load('diffuser_4_2.pt', map_location=DEVICE))
print("Generating sample...")
img = sample_minecraft(diffuser, vae)
img = (img * 255).astype(np.uint8)
Image.fromarray(img).save("generated_sample.png")
print("Saved to generated_sample.png")