pytorch如何知道某个Parameter是在哪一个Module中的创建的

pytorch如何知道某个Parameter是在哪一个Module中的创建的

在定位pytorch精度问题时,发现optimizer中某些Parameter值异常,想知道它属于哪个模块的.本文提供二种方法
1.全局搜索
2.在创建Parameter的地方加一个属性,写明所在的模块名,需要的时候直接获取

代码

python 复制代码
import torch
import sys
sys.setrecursionlimit(1000)
        
def search_recursive(var,stack,_id,depth):
    if var.__class__.__name__ in [
                                    "module","type","NoneType",
                                    "str","int","function","method-wrapper",
                                    "builtin_function_or_method",
                                    "method","_TensorMeta",
                                    "Tensor","method_descriptor",
                                    "bool","device","dtype",
                                    "getset_descriptor","layout",
                                    "wrapper_descriptor","property",
                                    "_ParameterMeta","mappingproxy",
                                    "Parameter","_abc_data","SourceFileLoader",
                                    "code","bytes","ABCMeta",
                                    "ForwardRef","ellipsis","TypeVar"
                                 ]:
        return False
    
    if isinstance(var,dict):
        for k,v in var.items():
            ret=search_recursive(v,stack,_id,depth+1)
            if ret:
                return ret
    elif isinstance(var,list) or isinstance(var,tuple):
        for i in var:
            ret=search_recursive(i,stack,_id,depth+1)
            if ret:
                return ret
    else:     
        if not var.__class__.__name__.startswith("_"):
            stack[depth]=var.__class__.__name__         
        for name in dir(var):
            try:
                obj=eval(f"var.{name}")
                if isinstance(obj,torch.nn.modules.linear.Linear) and id(obj.weight)==_id:
                    return stack[depth]
                ret=search_recursive(obj,stack,_id,depth+1)
                if ret:
                    return ret                 
            except:
                pass              
    return None

class MyModel(torch.nn.Module):
    def __init__(self):
        super(MyModel, self).__init__()
        self.mlp=torch.nn.Linear(5120,3850)
        self.mlp.weight.__setattr__("model_name","MyModel") #方法一:通过添加属性
    def forward(self, x):
        out=self.mlp(x)
        return out
class MyContainer(object):
    def __init__(self):
        self.obj=MyModel()    
    def get_param(self):
        return self.obj.mlp.weight
obj = MyContainer()
param_group={}
param_group["w0"]=obj.get_param()
param_array=[param_group,obj]

print("GetModelName By getattr:",getattr(param_group["w0"],"model_name"))
# 方法二:递归搜索全局变量
model_name=search_recursive(globals(),{},id(param_group["w0"]),0)
print("GetModelName By search_recursive:",model_name)
相关推荐
搬砖的小码农_Sky1 分钟前
AI Agent:如何处理Claude Code 最近版本(2026年更新)引入的模型上下文限制
人工智能·windows·ai·ai编程
企业数字化笔记11 分钟前
视频目标跟踪怎么选?SORT、DeepSORT、ByteTrack 的连续性、遮挡与计算成本对比
人工智能·目标跟踪·音视频
言乐616 分钟前
Python根据无法识别搜索词找出可能输入内容模型
开发语言·python·django·virtualenv·pygame
weixin_3077791321 分钟前
有限产能智能排产与动态重排智能体:从需求解构到技术实现
开发语言·人工智能·算法·架构
Ivanqhz26 分钟前
激活函数在 Transformer 中的作用及各种变体简述
java·linux·数据库·人工智能·深度学习
染指111027 分钟前
134.Agent-多Agent框架-LangChain多智能体
人工智能·中间件·langchain·agents
fkyyly1 小时前
企业 Agent 落地的两场战争:内部经营闭环与 ToB 外部规模赋能
人工智能·codeagent
看浪的路人1 小时前
第9讲:AI 应用混沌工程与容灾演练
人工智能
workflower1 小时前
AI system product quality model
大数据·人工智能·机器学习·云计算·无人机
众链网络1 小时前
从 GB/T 30225-2026 的「数据管理」章节,反推景区票务系统的数据层设计
人工智能