activation_swap.py(torch 与 mindspore 两个平台都有)在给模块施加 activation swap / 重计算 checkpoint 时,_check_and_mark_wrapped 会遍历该模块及其所有子模块,对每个“可包装的公开 callable 实例属性”做两件事:先校验它没有被标记过 _is_wrapped(重叠检测),再给它打上 _is_wrapped = True。
activation_swap.py
_check_and_mark_wrapped
_is_wrapped
_is_wrapped = True
判定“可包装 callable 属性”的 _iter_wrappable_callable_attrs 之前会把实例上任何非 Cell/Module 的公开 callable 都算进来,其中包含大量以实例属性形式持有的模块级共享函数,例如每层都写的 self.act = F.gelu、self.reshape = mint.reshape、self.cast = ops.cast。
_iter_wrappable_callable_attrs
self.act = F.gelu
self.reshape = mint.reshape
self.cast = ops.cast
这些是无状态、被多个模块以引用方式共享的同一个对象,把它纳入重叠跟踪有两个问题:
self.act
F.gelu
self.act._is_wrapped == True
该问题在重计算独立调度(对多个同级模块分别加 checkpoint)的场景下暴露。
_check_callable_attr_not_wrapped
ValueError: Callable 'function' is already wrapped. Wrapping overlapping module regions is not allowed.
该问题是怎么引起的?
activation_swap.py(torch 与 mindspore 两个平台都有)在给模块施加 activation swap / 重计算 checkpoint 时,_check_and_mark_wrapped会遍历该模块及其所有子模块,对每个“可包装的公开 callable 实例属性”做两件事:先校验它没有被标记过_is_wrapped(重叠检测),再给它打上_is_wrapped = True。判定“可包装 callable 属性”的
_iter_wrappable_callable_attrs之前会把实例上任何非 Cell/Module 的公开 callable 都算进来,其中包含大量以实例属性形式持有的模块级共享函数,例如每层都写的self.act = F.gelu、self.reshape = mint.reshape、self.cast = ops.cast。这些是无状态、被多个模块以引用方式共享的同一个对象,把它纳入重叠跟踪有两个问题:
_is_wrapped会污染全局对象(Python 普通函数允许写属性),影响程序里其它引用同一函数的位置。self.act(=F.gelu)标成 wrapped,再独立 wrap 同级的 B 层时,会看到self.act._is_wrapped == True,于是误判为区域重叠并抛出异常——而这两个区域其实并不重叠,只是共享了同一个无状态工具函数。该问题在重计算独立调度(对多个同级模块分别加 checkpoint)的场景下暴露。
重现步骤
self.act = F.gelu或self.reshape = mint.reshape。_check_and_mark_wrapped→_check_callable_attr_not_wrapped时,发现共享函数已被第一个模块标记,触发误报。报错信息