Views
No views yet

1import torch
2from torch import nn
3
4# Define generator architecture
5class Generator(nn.Module):
6 # ... (architecture code)
7
8# Load the model
9generator = Generator(latent_dim=100, channels=3)
10generator.load_state_dict(torch.load('model/generator_epoch_8800.pt'))
11generator.eval()
12
13# Generate images
14z = torch.randn(1, 100)
15with torch.no_grad():
16 fake_image = generator(z)
17 # Convert to range [0, 1] for display
18 fake_image = (fake_image + 1) / 2.0