Views
No views yet
1from diffusers import DDPMPipeline
2
3pipeline = DDPMPipeline.from_pretrained('your_hub_username/ddpm-fewshot_anime_face')
4image = pipeline().images[0]
5image.show() # To display the generated image
6
7#也可以使用下面的方法微调模型,我当前数据量较少
8from torchvision import transforms
9import torch
10from diffusers import DDPMPipeline,DDIMScheduler
11from datasets import load_dataset
12from torch.utils.data import DataLoader
13from tqdm import tqdm_notebook
14import torch.nn.functional as F
15from matplotlib import pyplot as plt
16
17path = "/Volumes/mac1/pytorchlearning/diff_learning/ddpm-fewshot_anime_face" #改为model_id
18image_pipe = DDPMPipeline.from_pretrained(path)
19scheduler = DDIMScheduler.from_pretrained(path)
20dataset = load_dataset("huggan/few-shot-anime-face",split="train")
21device = torch.device("mps" if torch.backends.mps.is_available() else "cpu")
22image_pipe.to(device)
23
24model_name = "ddpm-fewshot_anime_face"
25image_size = 256
26batch_size = 8
27preprocess = transforms.Compose([
28 transforms.Resize(image_size),
29 transforms.RandomHorizontalFlip(),
30 transforms.ToTensor(),
31 transforms.Normalize([0.5],[0.5])
32])
33
34def transform(examples):
35 images = [preprocess(image.convert("RGB")) for image in examples["image"]]
36 return {"images":images}
37
38dataset.set_transform(transform)
39train_loader = DataLoader(dataset,batch_size,shuffle=True)
40
41num_epochs = 20
42lr = 1e-5
43grad_accumulation_steps = 2
44optimizer = torch.optim.AdamW(image_pipe.unet.parameters(), lr=lr)
45losses = []
46for epoch in range(num_epochs):
47 for step,batch in tqdm_notebook(enumerate(train_loader), total=len(train_loader)):
48 clean_images = batch["images"].to(device)
49 noise = torch.randn(clean_images.shape).to(device)
50 bs = clean_images.shape[0]
51
52 time_steps = torch.randint(0,image_pipe.scheduler.num_train_timesteps,(bs,),device=device).long()
53 noise_image = image_pipe.scheduler.add_noise(clean_images,noise,time_steps)
54 noise_pred = image_pipe.unet(noise_image,time_steps,return_dict=False)[0]
55 loss = F.mse_loss(noise_pred,noise_image)
56
57 losses.append(loss.item())
58 loss.backward()
59 if (step + 1) % grad_accumulation_steps == 0:
60 optimizer.step()
61 optimizer.zero_grad()
62
63 print(f"Epoch {epoch} average loss: {sum(losses[-len(train_loader):])/len(train_loader)}")
64 image_pipe.save_pretrained(model_name)
65
66plt.plot(losses)