PyTorch中torchinfo.summary参数计数异常与recursive含义咨询
问题解析
1. 元组乘法构建隐藏层的问题
你用元组乘法(比如(nn.Linear(xxx, xxx),) * 2)生成2层隐藏层时,本质是创建了同一个Linear模块的多个引用,而非两个独立的Linear层。简单说,这两个“层”其实是同一个对象,共享一套参数,并不是各自拥有独立的参数。
这就导致torchinfo统计总参数时,只会计算这一个模块的参数一次,而不是两次。比如单个Linear层若有420个参数,你预期两个层贡献840个参数,但实际只算420个,加上输入输出层的参数后,总参数自然远低于预期,也就是你看到的501 vs 预期921的情况。
而显式元组相加(比如(nn.Linear(xxx, xxx),) + (nn.Linear(xxx, xxx),))是创建两个完全独立的Linear实例,每个层都有自己的参数,所以torchinfo能正确统计两次参数,总参数符合预期。
2. summary中(recursive)的含义
torchinfo里标注的(recursive),是用来提示你该模块实例在模型结构中被重复引用了。因为元组乘法生成的是同一个模块的多个引用,当torchinfo遍历模型的模块树时,会发现同一个模块对象已经被统计过一次,再次遇到时就会标注(recursive),告诉你这不是一个新创建的模块,而是之前出现过的模块的复用。
内容的提问来源于stack exchange,提问作者Rajesh Nakka
相关产品推荐
相关产品推荐

