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

TensorFlow GPU运行缓慢——疑为重追踪(retracing)问题如何解决?

解决TensorFlow中可变序列长度导致的重追踪性能问题

我正在开发Transformer序列转向量模型,使用Tesla V100-PCIE-16GB GPU,遇到性能问题:当模型处理未见过的序列长度时,训练步骤耗时显著增加。推测这是TensorFlow的图追踪算法导致的重追踪,但未收到任何重追踪警告。了解@tf.function(reduce_retracing=True)装饰器可缓解该问题,但需要从头重写MultiHeadAttention类并为每个函数添加装饰器,当前使用TensorFlow 2.17.0版本。

以下是复现代码:示例网络训练100个样本,前10批序列形状各异,每批耗时约1.5秒;后90批形状相似,每批仅耗时约0.002秒。请问该如何解决?

import datetime
import numpy as np
import tensorflow as tf
from tensorflow import keras

# 仅使用指定GPU
gpus = tf.config.list_physical_devices("GPU")
assert gpus, "未找到GPU设备"
gpu = gpus[1]
tf.config.set_visible_devices([gpu], "GPU")
tf.config.experimental.set_memory_growth(gpu, True)

nclasses = 3
d_model = 8

# 自定义数据生成器,生成随机样本
class DG(keras.utils.PyDataset):
    def __init__(self):
        self.nSamplesWithAlternatingShapes = 10
        super().__init__()
        self.last = datetime.datetime.now()

    def __len__(self):
        return 100

    def __getitem__(self, index):
        n = self.nSamplesWithAlternatingShapes
        if index < self.nSamplesWithAlternatingShapes:
            n = index
        x = np.random.random((32, n + 10, d_model))
        y = np.random.random((32, nclasses))
        y = np.argmax(y, axis=1)
        y = tf.keras.utils.to_categorical(y, num_classes=nclasses)
        now = datetime.datetime.now()
        print(f"\n耗时 {round((now-self.last).total_seconds(), ndigits=3)} 秒")
        self.last = now
        return tf.convert_to_tensor(x), y

# 测试层,包含MultiHeadAttention
class TestLayer(tf.keras.layers.Layer):
    def __init__(self):
        super().__init__()
        # 多头注意力层
        self.mha = tf.keras.layers.MultiHeadAttention(
            num_heads=1, value_dim=d_model, key_dim=d_model, dropout=0.3
        )
        self.dense = tf.keras.layers.Dense(nclasses, activation="softmax")

    def call(self, x):
        x = self.mha(x, x, x)
        pool = tf.keras.layers.GlobalAveragePooling1D(data_format="channels_last")(x)
        fin = self.dense(pool)
        return fin


if __name__ == "__main__":
    # 输入形状:(None, d_model),None表示可变序列长度
    input = tf.keras.layers.Input((None, d_model), name="test", dtype=np.float32)
    output = TestLayer()(input)
    model = tf.keras.models.Model(inputs=[input], outputs=output)
    model.compile(
        optimizer="adam",
        loss="categorical_crossentropy",
        metrics=[tf.keras.metrics.CategoricalAccuracy(name="OA")],
    )
    
    model.fit(DG(), shuffle=False)
    
    print("训练完成")

解决方案

1. 为自定义Layer的call方法添加reduce_retracing装饰器

无需重写整个MultiHeadAttention类,只需在自定义TestLayer的call方法上添加@tf.function(reduce_retracing=True),同时将层实例化移到__init__方法中(避免在call内重复创建层引发额外图变化):

class TestLayer(tf.keras.layers.Layer):
    def __init__(self):
        super().__init__()
        self.mha = tf.keras.layers.MultiHeadAttention(
            num_heads=1, value_dim=d_model, key_dim=d_model, dropout=0.3
        )
        self.dense = tf.keras.layers.Dense(nclasses, activation="softmax")
        # 将池化层移到__init__中实例化
        self.global_pool = tf.keras.layers.GlobalAveragePooling1D(data_format="channels_last")

    @tf.function(reduce_retracing=True)
    def call(self, x):
        x = self.mha(x, x, x)
        pool = self.global_pool(x)
        fin = self.dense(pool)
        return fin

2. 启用重追踪警告验证推测

在代码开头添加以下配置,确认重追踪行为是否发生:

# 启用详细的图追踪日志
tf.autograph.set_verbosity(3)
# 确保不使用 eager 执行
tf.config.run_functions_eagerly(False)

也可以设置环境变量export TF_CPP_MIN_LOG_LEVEL=0,查看完整的TensorFlow日志输出。

3. 固定序列长度(可选)

如果业务场景允许,可设置最大序列长度,对短序列进行padding,让模型仅追踪一次对应最大长度的计算图:

MAX_SEQ_LEN = 20  # 根据实际业务设置上限
input = tf.keras.layers.Input((MAX_SEQ_LEN, d_model), name="test", dtype=np.float32)

修改数据生成器的__getitem__方法,对序列进行padding:

x = np.random.random((32, n + 10, d_model))
if x.shape[1] < MAX_SEQ_LEN:
    pad_len = MAX_SEQ_LEN - x.shape[1]
    # 在序列维度进行padding
    x = np.pad(x, ((0,0), (0,pad_len), (0,0)), mode='constant')

4. 升级TensorFlow版本(可选)

TensorFlow 2.18及以上版本对Keras层的图优化有改进,可能原生减少MultiHeadAttention的重追踪情况,可尝试升级版本验证效果。


内容的提问来源于stack exchange,提问作者D. E.

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.17 06:57:03