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

如何用torchinfo生成GAN生成器及完整GAN的网络结构摘要

问题解答

1. 生成器完整摘要生成方法

完全可以用torchinfo为生成器输出和判别器同等详细度的结构摘要,你之前参数传错了。
torchinfo生成完整摘要的核心逻辑是执行一次真实的前向传播,通过hook捕获每一层的输入输出张量,才能计算出每层输出形状、乘加计算量、显存占用这类统计数据。你之前给生成器传output_size参数的方式是错误的,这个参数不触发完整前向追踪,只能遍历模型结构算出参数量,自然缺失其他字段。
生成器的输入是隐空间噪声,你只需要传入噪声的维度即可,参考你贴的生成器第一层Linear参数量(25856)推算,你的噪声维度为100,正确调用代码如下:

model = Generator()
batch_size = 32
noise_dim = 100 # 替换为你实际使用的隐向量维度
summary(model, input_size=(batch_size, noise_dim))

执行后就能得到和判别器完全一致的、包含所有统计字段的完整摘要。

2. 同时覆盖生成器+判别器的GAN整体摘要方案

存在两种成熟的实现方案,按需选择即可:

  • 方案一:封装端到端GAN模型(推荐,可获取全链路统计)
    写一个简单的包装类把生成器、判别器包含在内,forward方法实现从噪声输入到判别器输出的完整推理流程,直接传入噪声维度即可一次性生成整个GAN的完整结构摘要,示例代码:
    import torch.nn as nn
    from torchinfo import summary
    
    class FullGAN(nn.Module):
        def __init__(self, generator, discriminator):
            super().__init__()
            self.generator = generator
            self.discriminator = discriminator
        
        def forward(self, noise):
            fake_images = self.generator(noise)
            d_logits = self.discriminator(fake_images)
            return d_logits
    
    # 调用示例
    batch_size = 32
    noise_dim = 100
    gan_model = FullGAN(Generator(), Discriminator())
    summary(gan_model, input_size=(batch_size, noise_dim))
    
    该方式输出的摘要会按执行顺序列出生成器、判别器的所有层,同时自动计算整个GAN端到端的总参数量、总乘加计算量、显存占用等完整统计信息。
  • 方案二:分别生成两个子网络的摘要
    如果你需要单独查看两个子网络的独立统计数据,不需要端到端链路信息,直接分别对生成器、判别器调用summary,传入各自对应的输入尺寸即可,最后把两份结果拼接使用。

内容的提问来源于stack exchange,提问作者Sterik

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 15:51:27