Generator Tích chập
Định nghĩa một generator tích chập theo các nguyên tắc DCGAN đã thảo luận trong video trước.
torch.nn đã được nhập sẵn dưới tên nn để bạn tiện sử dụng. Ngoài ra, hàm tùy chỉnh dc_gen_block() cũng đã được cung cấp, hàm này trả về một khối gồm tích chập chuyển vị, chuẩn hóa theo batch và kích hoạt ReLU. Hàm này đóng vai trò là thành phần nền tảng để xây dựng generator tích chập. Bạn có thể xem định nghĩa của dc_gen_block() bên dưới.
def dc_gen_block(in_dim, out_dim, kernel_size, stride):
return nn.Sequential(
nn.ConvTranspose2d(in_dim, out_dim, kernel_size, stride=stride),
nn.BatchNorm2d(out_dim),
nn.ReLU()
)
Bài tập này là một phần của khóa học
Deep Learning cho Ảnh với PyTorch
Hướng dẫn bài tập
- Thêm khối generator cuối, ánh xạ kích thước của các feature map thành
256. - Thêm một tích chập chuyển vị với kích thước đầu ra là
3. - Thêm hàm kích hoạt tanh.
Bài tập tương tác thực hành trực tiếp
Hãy thử làm bài tập này bằng cách hoàn thành đoạn mã mẫu này.
class DCGenerator(nn.Module):
def __init__(self, in_dim, kernel_size=4, stride=2):
super(DCGenerator, self).__init__()
self.in_dim = in_dim
self.gen = nn.Sequential(
dc_gen_block(in_dim, 1024, kernel_size, stride),
dc_gen_block(1024, 512, kernel_size, stride),
# Add last generator block
____,
# Add transposed convolution
____(____, ____, kernel_size, stride=stride),
# Add tanh activation
____
)
def forward(self, x):
x = x.view(len(x), self.in_dim, 1, 1)
return self.gen(x)