You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何深度复制被@tf.function装饰的TensorFlow函数?

解决@tf.function装饰函数的deepcopy问题

问题描述

使用copy.deepcopy复制被@tf.function装饰的函数时,函数未初始化前可正常复制,但调用函数完成初始化后,会抛出TypeError: can't pickle _thread.RLock objects错误。

复现代码

import tensorflow as tf
import copy

# 定义被装饰的函数
@tf.function
def foo(x): return x * 2

# 未初始化时可正常执行deepcopy
print(copy.deepcopy(foo))  # <tensorflow.python.eager.def_function.Function object at ...>

# 调用函数完成初始化
foo(tf.constant([3.]))

# 此时执行deepcopy会失败
copy.deepcopy(foo)  # 抛出TypeError: can't pickle _thread.RLock objects

问题根源:初始化后的Function对象包含_stateful_fn、_concrete_stateful_fn等嵌套属性,这些属性内部存在无法被pickle的_thread.RLock锁对象;虽然Function类自带的__getstate__方法会忽略顶层锁,但深拷贝会递归处理嵌套属性,导致锁对象暴露引发错误。

可行解决方案

方案1:基于原始函数重新创建装饰实例

如果接受复制后的函数首次调用时重新编译,直接基于未装饰的原始Python函数重新应用@tf.function即可,完全绕过pickle问题:

# 先保留原始未装饰的函数
def _foo(x): return x * 2

# 生成装饰后的原函数
foo = tf.function(_foo)

# 初始化后,复制时直接重新装饰原始函数
foo_copy = tf.function(_foo)

该方式简单可靠,复制后的函数首次调用会重新编译,但原函数的编译结果不受影响。

方案2:自定义深拷贝逻辑,清理不可pickle属性

通过手动清理Function对象中的不可pickle嵌套属性,生成新的未初始化实例:

import copy
from tensorflow.python.eager.def_function import Function

def custom_deepcopy(func):
    if not isinstance(func, Function):
        return copy.deepcopy(func)
    
    # 利用原类的__getstate__过滤顶层不可pickle属性
    state = func.__getstate__()
    
    # 移除已编译的嵌套属性,重置为未初始化状态
    for key in ['_stateful_fn', '_stateless_fn', '_concrete_stateful_fn', '_lifted_initializer_graph']:
        state.pop(key, None)
    
    # 创建新的Function实例并恢复状态
    new_func = Function(func.python_function, func._func_graph_spec, func._name)
    new_func.__setstate__(state)
    return new_func

# 使用示例
foo(tf.constant([3.]))  # 初始化原函数
foo_copy = custom_deepcopy(foo)
print(foo_copy(tf.constant([5.])))  # 首次调用重新编译,输出tf.Tensor([10.], shape=(1,), dtype=float32)

此方法保留原函数的基础配置,复制后的函数无需重新定义原始Python函数,仅首次调用时重新编译。

方案3:用pickle替代deepcopy(仅序列化场景)

如果只是为了序列化操作,直接使用pickle模块即可——Function的__getstate__已处理顶层不可pickle属性,且pickle不会像deepcopy那样递归处理嵌套属性:

import pickle

# 初始化后的函数可正常序列化
pickled_data = pickle.dumps(foo)
foo_copy = pickle.loads(pickled_data)

print(foo_copy(tf.constant([4.])))  # 输出tf.Tensor([8.], shape=(1,), dtype=float32)

注意:pickle加载后的函数首次调用会重新编译,但原函数的编译状态不受影响。

说明

  • 以上方案均满足“复制后的版本可重新初始化,不影响原函数编译状态”的需求;
  • 优先推荐方案1(最简洁)或方案3(仅序列化场景),方案2适合需要保留原Function实例配置的场景。

内容的提问来源于stack exchange,提问作者Uri Granta

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.21 18:42:40