如何提取PyTorch网络容器的类名
如何提取PyTorch网络容器的类名
嘿,这个问题我刚好折腾过,其实不用费劲去手动拆解type()返回的完整类路径,有几个简单直接的方法能拿到Sequential这类容器的类名:
最直观的方法:用
.__class__.__name__
直接调用模型容器的这个属性就能得到纯类名字符串,比如针对你的例子:print(p.__class__.__name__)运行后会直接输出
Sequential,是不是很省心?原理也很简单:__class__会返回模块对应的类对象,而__name__就是这个类的名称字符串。另一种等价写法:
type(p).__name__
如果你习惯用type()函数,也可以直接取它的__name__属性,效果完全一样:print(type(p).__name__)同样会输出
Sequential,比你手动去切割<class 'torch.nn.modules.container.Sequential'>这种字符串要高效得多。
如果你想把容器名和子模块一起做自定义的格式化打印,还可以把这两种方法和你原来的遍历代码结合起来,比如:
# 先打印容器类名 print(f"{p.__class__.__name__}(") # 遍历子模块并格式化输出 for name, module in p.named_children(): print(f' ({name}): {module}') print(")")
这样输出的格式就和原生print(p)差不多,但你可以根据自己的需求调整缩进、排版,灵活性更高。
另外多说一句,这两种方法对PyTorch里其他容器(比如ModuleList、ModuleDict)也完全适用,通用性拉满~
备注:内容来源于stack exchange,提问作者JKomp
相关产品推荐
相关产品推荐

