本文将开始逐步复现代码,从 face render 开始。
2024.11.01 更新 encoder , 即论文中的 和
, 二者结构一模一样,但是分别用于对不同图像进行编码,前者用于编码人脸参考图,后者用于编码所有的 mesh。
from torch import nn
import torch.nn.functional as F
import torch
import torch.nn.utils.spectral_norm as spectral_norm
import math
import warnings
class SameBlock2d(nn.Module):
"""
Simple block, preserve spatial resolution.
"""
def __init__(self, in_features, out_features, groups=1, kernel_size=3, padding=1, lrelu=False):
super(SameBlock2d, self).__init__()
self.conv = nn.Conv2d(in_channels=in_features, out_channels=out_features, kernel_size=kernel_size, padding=padding, groups=groups)
self.norm = nn.BatchNorm2d(out_features, affine=True)
if lrelu:
self.ac = nn.LeakyReLU()
else:
self.ac = nn.ReLU()
def forward(self, x):
out = self.conv(x)
out = self.norm(out)
out = self.ac(out)
return out
class DownBlock2d(nn.Module):
"""
Downsampling block for use in encoder.
"""
def __init__(self, in_features, out_features, kernel_size=3, padding=1, groups=1):
super(DownBlock2d, self).__init__()
self.conv = nn.Conv2d(in_channels=in_features, out_channels=out_features, kernel_size=kernel_size, padding=padding, groups=groups)
self.norm = nn.BatchNorm2d(out_features, affine=True)
self.pool = nn.AvgPool2d(kernel_size=(2, 2))
def forward(self, x):
out = self.conv(x)
out = self.norm(out)
out = F.relu(out)
out = self.pool(out)
return out
class AppearanceFeatureExtractor2D(nn.Module):
def __init__(self, image_channel, block_expansion, num_down_blocks, max_features):
super(AppearanceFeatureExtractor2D, self).__init__()
self.image_channel = image_channel
self.block_expansion = block_expansion
self.num_down_blocks = num_down_blocks
self.max_features = max_features
self.first = SameBlock2d(
image_channel, block_expansion, kernel_size=(3, 3), padding=(1, 1)
)
down_blocks = []
for i in range(num_down_blocks):
in_features = min(max_features, block_expansion * (2**i))
out_features = min(max_features, block_expansion * (2 ** (i + 1)))
down_blocks.append(
DownBlock2d(
in_features, out_features, kernel_size=(3, 3), padding=(1, 1)
)
)
self.down_blocks = nn.ModuleList(down_blocks)
def forward(self, source_image):
out = self.first(source_image) # Bx3x256x256 -> Bx64x256x256
for i in range(len(self.down_blocks)):
out = self.down_blocks[i](out)
return out
根据此代码,假设输入是 ,那么输出就是
。

1410

被折叠的 条评论
为什么被折叠?



