如何将flax.linen.Module转换为torch.nn.Module?核心层替换答疑
Flax转PyTorch:ScoreNet实现迁移指南
核心替换方案说明
1. flax.linen.Dense → torch.nn.Linear
Flax的Dense(output_dim)等价于PyTorch的Linear(in_features, out_features),参数对应规则:
out_features直接对应Flax的output_dimin_features需根据输入张量的最后一维维度确定,比如时间嵌入分支中,GaussianFourierProjection输出维度为embed_dim,因此此处Linear的in_features设为embed_dim
2. flax.linen.Conv → torch.nn.Conv2d
Flax的Conv(output_channels, kernel_size, stride, padding, ...)对应PyTorch的Conv2d(in_channels, out_channels, kernel_size, stride, padding, dilation, ...),关键参数映射:
out_channels对应Flax的第一个参数(输出通道数)in_channels根据前一层输出的通道数确定(比如输入为单通道图像时,第一个Conv2d的in_channels=1)- Flax的
padding='VALID'→ PyTorch的padding=0 - Flax的
padding=((a,b),(c,d))→ PyTorch的padding=(c,d)(维度顺序一致,均为(H,W)) - Flax的
input_dilation→ PyTorch的dilation(两者均实现卷积输入的膨胀效果)
3. 自定义Dense类迁移
Flax中自定义Dense类的作用是将全连接层输出reshape为[batch, 1, 1, output_dim],方便与卷积特征图相加。PyTorch中可通过Linear层后添加unsqueeze生成两个空间维度实现该逻辑。
完整PyTorch实现代码
import torch import torch.nn as nn import torch.nn.functional as F from typing import Any, Tuple class GaussianFourierProjection(nn.Module): """Gaussian random features for encoding time steps.""" def __init__(self, embed_dim: int, scale: float = 30.): super().__init__() # 初始化固定权重,不参与训练 self.W = nn.Parameter(torch.randn(embed_dim // 2) * scale, requires_grad=False) def forward(self, x): x_proj = x[:, None] * self.W[None, :] * 2 * torch.pi return torch.cat([torch.sin(x_proj), torch.cos(x_proj)], dim=-1) class Dense(nn.Module): """A fully connected layer that reshapes outputs to feature maps.""" def __init__(self, input_dim: int, output_dim: int): super().__init__() self.linear = nn.Linear(input_dim, output_dim) def forward(self, x): # 输出形状: [batch, 1, 1, output_dim] return self.linear(x).unsqueeze(1).unsqueeze(1) class ScoreNet(nn.Module): """A time-dependent score-based model built upon U-Net architecture.""" def __init__(self, marginal_prob_std: Any, channels: Tuple[int] = (32, 64, 128, 256), embed_dim: int = 256): super().__init__() self.marginal_prob_std = marginal_prob_std self.channels = channels self.embed_dim = embed_dim # 时间嵌入分支 self.fourier_proj = GaussianFourierProjection(embed_dim=embed_dim) self.dense_embed = nn.Linear(embed_dim, embed_dim) # 编码路径 self.conv1 = nn.Conv2d(1, channels[0], kernel_size=(3,3), stride=(1,1), padding=0, bias=False) self.dense1 = Dense(embed_dim, channels[0]) self.norm1 = nn.GroupNorm(4, channels[0]) self.conv2 = nn.Conv2d(channels[0], channels[1], kernel_size=(3,3), stride=(2,2), padding=0, bias=False) self.dense2 = Dense(embed_dim, channels[1]) self.norm2 = nn.GroupNorm(1, channels[1]) self.conv3 = nn.Conv2d(channels[1], channels[2], kernel_size=(3,3), stride=(2,2), padding=0, bias=False) self.dense3 = Dense(embed_dim, channels[2]) self.norm3 = nn.GroupNorm(1, channels[2]) self.conv4 = nn.Conv2d(channels[2], channels[3], kernel_size=(3,3), stride=(2,2), padding=0, bias=False) self.dense4 = Dense(embed_dim, channels[3]) self.norm4 = nn.GroupNorm(1, channels[3]) # 解码路径 self.deconv4 = nn.Conv2d(channels[3], channels[2], kernel_size=(3,3), stride=(1,1), padding=(2,2), dilation=(2,2), bias=False) self.dense_deconv4 = Dense(embed_dim, channels[2]) self.norm_deconv4 = nn.GroupNorm(1, channels[2]) self.deconv3 = nn.Conv2d(channels[2]+channels[2], channels[1], kernel_size=(3,3), stride=(1,1), padding=(2,3), dilation=(2,2), bias=False) self.dense_deconv3 = Dense(embed_dim, channels[1]) self.norm_deconv3 = nn.GroupNorm(1, channels[1]) self.deconv2 = nn.Conv2d(channels[1]+channels[1], channels[0], kernel_size=(3,3), stride=(1,1), padding=(2,3), dilation=(2,2), bias=False) self.dense_deconv2 = Dense(embed_dim, channels[0]) self.norm_deconv2 = nn.GroupNorm(1, channels[0]) self.deconv1 = nn.Conv2d(channels[0]+channels[0], 1, kernel_size=(3,3), stride=(1,1), padding=(2,2)) def forward(self, x, t): act = F.silu # PyTorch的silu对应Flax的swish # 时间嵌入 embed = act(self.dense_embed(self.fourier_proj(t))) # 编码路径 h1 = self.conv1(x) h1 += self.dense1(embed) h1 = self.norm1(h1) h1 = act(h1) h2 = self.conv2(h1) h2 += self.dense2(embed) h2 = self.norm2(h2) h2 = act(h2) h3 = self.conv3(h2) h3 += self.dense3(embed) h3 = self.norm3(h3) h3 = act(h3) h4 = self.conv4(h3) h4 += self.dense4(embed) h4 = self.norm4(h4) h4 = act(h4) # 解码路径 h = self.deconv4(h4) h += self.dense_deconv4(embed) h = self.norm_deconv4(h) h = act(h) h = torch.cat([h, h3], dim=1) h = self.deconv3(h) h += self.dense_deconv3(embed) h = self.norm_deconv3(h) h = act(h) h = torch.cat([h, h2], dim=1) h = self.deconv2(h) h += self.dense_deconv2(embed) h = self.norm_deconv2(h) h = act(h) h = torch.cat([h, h1], dim=1) h = self.deconv1(h) # 输出归一化 h = h / self.marginal_prob_std(t)[:, None, None, None] return h
内容的提问来源于stack exchange,提问作者0xbadf00d
相关产品推荐
相关产品推荐

