使用tensorflow-directml训练时model.fit()首个epoch中途停滞的解决方法
问题现象
- 初始调用
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版本依赖该属性,就会触发此错误。
解决办法
- 升级Python版本(推荐)
将Python升级至3.8或更高版本,匹配TensorFlow-DirectML的依赖要求。 - 降级TensorFlow-DirectML
若无法升级Python,安装兼容旧版本Python的TensorFlow-DirectML 2.10.x:pip uninstall tensorflow-directml -y pip install tensorflow-directml==2.10.0 - 自定义训练循环规避
用自定义训练循环替代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

