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

使用tensorflow-directml训练时model.fit()首个epoch中途停滞的解决方法

解决AMD RX580 + TensorFlow-DirectML训练卡住及AutoGraph警告问题

问题现象

  • 初始调用model.fit()时卡在第1个epoch无进展
  • 触发AutoGraph转换失败警告:

    Epoch 1/15
    WARNING:tensorflow:AutoGraph could not transform <function Model.make_train_function..train_function at 0x000002063B21A840> and will run it as-is.
    Please report this to the TensorFlow team. When filing the bug, set the verbosity to 10 (on Linux, export AUTOGRAPH_VERBOSITY=10) and attach the full output.
    Cause: 'arguments' object has no attribute 'posonlyargs'
    To silence this warning, decorate the function with @tf.autograph.experimental.do_not_convert

  • 调整参数后训练推进至41/758后停滞,脚本持续运行无输出

一、处理AutoGraph警告

原因

该警告源于Python版本与TensorFlow-DirectML版本不兼容:posonlyargs是Python 3.8+新增属性,若使用Python 3.7及以下版本,同时TensorFlow-DirectML版本依赖该属性,就会触发此错误。

解决办法

  1. 升级Python版本(推荐)
    将Python升级至3.8或更高版本,匹配TensorFlow-DirectML的依赖要求。
  2. 降级TensorFlow-DirectML
    若无法升级Python,安装兼容旧版本Python的TensorFlow-DirectML 2.10.x:
    pip uninstall tensorflow-directml -y
    pip install tensorflow-directml==2.10.0
    
  3. 自定义训练循环规避
    用自定义训练循环替代model.fit(),给训练步骤函数添加@tf.autograph.experimental.do_not_convert装饰器:
    import tensorflow as tf
    
    # 替换为你的模型对应的损失函数、优化器
    loss_fn = tf.keras.losses.CategoricalCrossentropy()
    optimizer = tf.keras.optimizers.Adam()
    train_loss = tf.keras.metrics.Mean(name='train_loss')
    train_accuracy = tf.keras.metrics.CategoricalAccuracy(name='train_accuracy')
    val_loss = tf.keras.metrics.Mean(name='val_loss')
    val_accuracy = tf.keras.metrics.CategoricalAccuracy(name='val_accuracy')
    
    @tf.autograph.experimental.do_not_convert
    def train_step(x, y):
        with tf.GradientTape() as tape:
            predictions = model(x, training=True)
            loss = loss_fn(y, predictions)
        gradients = tape.gradient(loss, model.trainable_variables)
        optimizer.apply_gradients(zip(gradients, model.trainable_variables))
        train_loss.update_state(loss)
        train_accuracy.update_state(y, predictions)
    
    @tf.autograph.experimental.do_not_convert
    def val_step(x, y):
        predictions = model(x, training=False)
        loss = loss_fn(y, predictions)
        val_loss.update_state(loss)
        val_accuracy.update_state(y, predictions)
    
    # 执行训练循环
    epochs = 15
    for epoch in range(epochs):
        train_loss.reset_states()
        train_accuracy.reset_states()
        val_loss.reset_states()
        val_accuracy.reset_states()
    
        # 训练步骤
        for step, (x_train, y_train) in enumerate(n_train_ds):
            train_step(x_train, y_train)
            if step % 10 == 0:
                print(f"Epoch {epoch+1}/{epochs}, Step {step}/{len(n_train_ds)}, Loss: {train_loss.result():.4f}")
    
        # 验证步骤
        for x_val, y_val in n_test_ds:
            val_step(x_val, y_val)
    
        print(f"Epoch {epoch+1}/{epochs} - Train Loss: {train_loss.result():.4f}, Train Acc: {train_accuracy.result():.4f} - Val Loss: {val_loss.result():.4f}, Val Acc: {val_accuracy.result():.4f}")
    

二、解决训练停滞问题

1. 优化数据集加载

确保数据集添加缓存和预取操作,避免CPU喂数据速度跟不上GPU:

n_train_ds = n_train_ds.cache().prefetch(tf.data.AUTOTUNE)
n_test_ds = n_test_ds.cache().prefetch(tf.data.AUTOTUNE)

同时检查batch size:RX580显存为8GB,若batch size超过32可能导致显存不足,尝试调小至16或8。

2. 移除手动设备指定

删除with tensorflow.device('/device:GPU:0')代码块,让TensorFlow自动分配设备。手动指定可能导致部分操作无法在GPU执行, fallback到CPU后出现阻塞。

3. 排查回调函数

临时移除checkpoint_callback和early_stop,确认是否是回调函数导致的停滞(如checkpoint保存路径无写入权限、early stop逻辑异常)。

4. 更新TensorFlow-DirectML版本

安装最新稳定版TensorFlow-DirectML:

pip uninstall tensorflow-directml -y
pip install tensorflow-directml

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.10 19:01:31