如何用字符串内容覆盖Python中的函数/方法?
问题:如何用修改后的AST代码覆盖PyTorch模型的forward方法?
问题背景
我想用一个第三方库处理CNN模型,但发现部分模型里的函数和这个库不兼容。排查后发现是因为部分层里用了+=运算符,这个运算符无法被该库处理,所以我需要替换所有这类代码。
已完成操作
我用inspect模块提取了有问题的函数字符串,再通过ast模块生成抽象语法树(AST),接着用NodeTransformer类修改了AST中的AugAssign(对应+=)节点,把它替换成符合要求的代码,而且已经把修改后的AST转回成字符串了。
当前困境
我试了好几种方法想把修改后的字符串转成可执行函数,用来覆盖原模型的forward方法,但都失败了。我的代码如下:
import torchvision import torch import inspect import textwrap import ast from pprint import pprint import torch.nn as nn # 待处理的模型 model = torchvision.models.resnet18() # 提取nn模块下的所有类 classes = [ x[0] for x in inspect.getmembers(nn, inspect.isclass) ] # 修改AST的转换器类 class Trasformer(ast.NodeTransformer): def visit_AugAssign(self,node): return (ast.Assign(targets=[ast.Name(id='out', ctx=ast.Store())], value=ast.Call(func=ast.Attribute(value=ast.Name(id='torch', ctx=ast.Load()), attr='stack', ctx=ast.Load()), args=[ast.List(elts=[ast.Name(id='identity', ctx=ast.Load()), ast.Name(id='out', ctx=ast.Load())], ctx=ast.Load())], keywords=[ast.keyword(arg='dim', value=ast.UnaryOp(op=ast.USub(), operand=ast.Constant(value=1)))])), ast.Assign(targets=[ast.Name(id='out', ctx=ast.Store())], value=ast.Call(func=ast.Attribute(value=ast.Name(id='self', ctx=ast.Load()), attr='canonizer_sum', ctx=ast.Load()), args=[ast.Name(id='out', ctx=ast.Load())], keywords=[]))) # 递归检查模型模块 def recursive(model): children = list(model.children()) print(f"Checking {model.__class__.__name__}\t{len(children)}") if len(children) == 0: pass else: if not model.__class__.__name__ in classes: print(f"!!{model.__class__.__name__}") print(textwrap.dedent(inspect.getsource(model.forward))) # 修改函数代码 forward_code = textwrap.dedent(inspect.getsource(model.forward)) tree = ast.parse(forward_code) Trasformer().visit(tree) tree = ast.fix_missing_locations(tree) forward_code = ast.unparse(tree) # 尝试覆盖forward方法,但失败 model.forward = exec(compile(tree, filename='test', mode='exec')) print("new string:\n") print(textwrap.dedent(inspect.getsource(model.forward))) for module in children: recursive(module) recursive(model) quit()
我的疑问
能不能用字符串内容覆盖函数/方法?我试过直接赋值、写入文件再加载、用exec/eval/compile,但都没成功。请问怎么用修改后的AST或者它生成的字符串覆盖原函数?Python标准库有没有可用的模块或方法?
解决方案
要实现用修改后的代码覆盖原forward方法,核心是把编译后的代码转换成能绑定到实例的函数——你之前的问题在于exec(compile(...))不会返回函数对象,得在执行时提取函数并手动绑定到实例上。
关键修改步骤
- 编译AST并提取函数:执行编译后的AST时,指定一个局部命名空间,用来捕获生成的
forward函数。 - 绑定函数为实例方法:用
types.MethodType把普通函数绑定成模型实例的方法,确保self参数能正确指向实例。
修改后的关键代码片段:
# 修改函数代码部分 forward_code = textwrap.dedent(inspect.getsource(model.forward)) tree = ast.parse(forward_code) Trasformer().visit(tree) ast.fix_missing_locations(tree) # 编译AST,在局部命名空间中执行 local_ns = {} exec(compile(tree, filename='__modified_forward__', mode='exec'), globals(), local_ns) # 从命名空间中取出修改后的forward函数 modified_forward = local_ns['forward'] # 绑定函数到模型实例 import types model.forward = types.MethodType(modified_forward, model) # 验证修改结果 print("修改后的forward代码:") print(textwrap.dedent(inspect.getsource(model.forward)))
额外注意事项
- AST结构合法性:你的
Transformer返回两个Assign节点,要确保这些节点被正确插入到原AST的body中,避免语法错误。可以在visit_AugAssign中返回一个节点列表,或者手动处理父节点的子元素。 - 实例属性依赖:如果修改后的代码用到
self.canonizer_sum这类实例属性,要确保目标模型实例已经初始化了这些属性,否则执行时会抛出AttributeError。 - 命名空间一致性:原
forward方法依赖的全局或局部变量,要保证在新函数的命名空间中能正常访问,必要时可以把相关变量传入exec的命名空间参数。
内容的提问来源于stack exchange,提问作者balticFoo
相关产品推荐
相关产品推荐

