如何在自定义BertClassifier模型上正确运行torchinfo
报错根因
- torchinfo仅传入输入形状元组时,默认生成
FloatTensor类型的测试输入做前向传播,但BERT的Embedding层要求input_ids必须是Int/Long类型的整数张量,类型不匹配直接触发RuntimeError。 - 传入的attention_mask形状配置错误:attention_mask和input_ids形状一致,均为
(batch_size, seq_len),配置里写的(4, 1, 512)多了一维,即使类型正确后续也会触发形状不匹配错误。
正确调用方式
两种可行配置方案二选一即可:
- 手动构造符合要求的测试输入传入,完全控制张量类型、形状
import torch from torchinfo import summary model = BertClassifier() summary( model, input_data=[ # input_id:整数类型,形状(批大小, 序列长度) torch.randint(low=0, high=model.bert.config.vocab_size, size=(4, 512), dtype=torch.long), # attention_mask:和input_id形状一致,整数类型 torch.ones(size=(4, 512), dtype=torch.long) ] )
- 保留形状传参方式,额外显式指定输入数据类型
summary( BertClassifier(), input_size=[(4, 512), (4, 512)], # 修正attention_mask形状,删除多余的1维度 dtypes=[torch.long, torch.long] # 两个输入均指定为Long类型,匹配Embedding层要求 )
如果模型部署在CUDA设备上,只需在
summary参数中追加device="cuda",torchinfo会自动将测试输入分配到对应设备,无需手动处理张量设备迁移。
内容的提问来源于stack exchange,提问作者user3668129
相关产品推荐
相关产品推荐

