如何在Python节点级别追踪代码执行及中间运算结果?
使用Python AST模块追踪代码执行中间结果完全可行
针对你提出的需求,通过AST模块解析并修改代码结构,插入自定义追踪逻辑,就能精准捕获示例代码中各中间步骤的输出形状与类型。
实现思路
核心是通过ast.NodeTransformer遍历代码的抽象语法树,对目标节点(np.reshape调用、@矩阵乘法、变量A的赋值)进行替换或扩展:
- 为
np.reshape调用添加包装逻辑:执行原函数后,打印结果的形状和类型,再返回结果继续参与后续运算。 - 为
@矩阵乘法添加包装逻辑:执行运算后打印结果信息,返回结果供赋值使用。 - 在变量A的赋值语句后,添加打印A最终形状和类型的代码。
具体代码实现
import ast import numpy as np def add_trace_to_code(code): # 解析原始代码为AST tree = ast.parse(code) class TraceTransformer(ast.NodeTransformer): def visit_Call(self, node): # 识别并包装np.reshape调用 if (isinstance(node.func, ast.Attribute) and node.func.attr == 'reshape' and isinstance(node.func.value, ast.Name) and node.func.value.id == 'np'): # 替换为trace_call包装调用 return ast.Call( func=ast.Name(id='trace_call', ctx=ast.Load()), args=[self.generic_visit(node), ast.Constant(value="np.reshape执行后")], keywords=[] ) return self.generic_visit(node) def visit_BinOp(self, node): # 识别并包装@矩阵乘法 if isinstance(node.op, ast.MatMult): return ast.Call( func=ast.Name(id='trace_binop', ctx=ast.Load()), args=[ self.generic_visit(node.left), self.generic_visit(node.right), ast.Constant(value="@运算执行后") ], keywords=[] ) return self.generic_visit(node) def visit_Assign(self, node): # 为变量A的赋值添加后续追踪 if len(node.targets) == 1 and isinstance(node.targets[0], ast.Name) and node.targets[0].id == 'A': processed_value = self.generic_visit(node.value) # 生成原始赋值语句 assign_stmt = ast.Assign(targets=node.targets, value=processed_value) # 生成打印A信息的语句 print_stmt = ast.Expr( value=ast.Call( func=ast.Name(id='print', ctx=ast.Load()), args=[ ast.Constant(value="赋值后变量A:形状="), ast.Attribute(value=ast.Name(id='A', ctx=ast.Load()), attr='shape', ctx=ast.Load()), ast.Constant(value=",类型="), ast.Call(func=ast.Name(id='type', ctx=ast.Load()), args=[ast.Name(id='A', ctx=ast.Load())], keywords=[]) ], keywords=[] ) ) # 返回赋值+打印的语句列表 return [assign_stmt, print_stmt] return self.generic_visit(node) # 添加辅助追踪函数到AST头部 helper_functions = """ def trace_call(expr, desc): res = expr print(f"{desc}:形状={res.shape},类型={type(res)}") return res def trace_binop(left, right, desc): res = left @ right print(f"{desc}:形状={res.shape},类型={type(res)}") return res """ helper_tree = ast.parse(helper_functions) tree.body = helper_tree.body + tree.body # 转换AST并修复位置信息 transformer = TraceTransformer() transformed_tree = transformer.visit(tree) ast.fix_missing_locations(transformed_tree) return transformed_tree # 测试示例代码 sample_code = """ import numpy as np R, Q, P = 2, 3, 4 A = np.random.rand(R, Q*P) B = np.random.rand(P, Q) A = np.reshape(A, (R, Q, 1, P)) @ B """ # 生成带追踪逻辑的AST并执行 traced_ast = add_trace_to_code(sample_code) exec(compile(traced_ast, filename='<ast_trace>', mode='exec'))
执行效果
运行上述代码后,控制台会输出类似如下内容:
np.reshape执行后:形状=(2, 3, 1, 4),类型=<class 'numpy.ndarray'> @运算执行后:形状=(2, 3, 1, 3),类型=<class 'numpy.ndarray'> 赋值后变量A:形状=(2, 3, 1, 3),类型=<class 'numpy.ndarray'>
完全覆盖了你需要的三个追踪目标。
内容的提问来源于stack exchange,提问作者Oscar
相关产品推荐
相关产品推荐

