You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何将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_dim
  • in_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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.26 16:05:33