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

使用TensorFlow-DirectML训练模型时遇形状不兼容错误求助

AMD GPU训练TensorFlow模型报错:形状不兼容[32] vs [32,528]

我在AMD GPU上使用基于TensorFlow 1.15.8的tensorflow-directml框架训练模型,CPU上用最新版TensorFlow可正常训练,但GPU训练时报错,提示形状不兼容:[32] vs [32,528],不清楚32的来源,模型输出类别仅为7类。

模型结构

Model: "sequential"
_________________________________________________________________
Layer (type)                 Output Shape              Param #   
=================================================================
NASNet (Model)               (None, 7, 7, 1056)        4269716   
_________________________________________________________________
global_average_pooling2d (Gl (None, 1056)              0         
_________________________________________________________________
batch_normalization (BatchNo (None, 1056)              4224      
_________________________________________________________________
reshape (Reshape)            (None, None, 1)           0         
_________________________________________________________________
average_pooling1d (AveragePo (None, None, 1)           0         
_________________________________________________________________
dropout (Dropout)            (None, None, 1)           0         
_________________________________________________________________
dense (Dense)                (None, None, 128)         256       
_________________________________________________________________
dropout_1 (Dropout)          (None, None, 128)         0         
_________________________________________________________________
dense_1 (Dense)              (None, None, 7)           903       
=================================================================
Total params: 4,275,099
Trainable params: 3,271
Non-trainable params: 4,271,828
_________________________________________________________________

训练代码

learning_rate = 0.001
optimizer = tensorflow.keras.optimizers.Adam(learning_rate=learning_rate)
loss = tensorflow.keras.losses.CategoricalCrossentropy(from_logits=False)
model.compile(optimizer=optimizer, loss=loss, metrics=['accuracy'])

with tensorflow.device('/device:DML:0'):
    history = model.fit(m_train_ds,
                        epochs=15,
                        steps_per_epoch=len(m_train_ds), #steps = 758
                        validation_data=m_test_ds,
                        validation_steps=len(m_test_ds), #steps = 190
                        callbacks=[checkpoint_callback, early_stop],
                        verbose=1,
                        class_weight=m_class_weights
                       )
   紧密Well Quick主导学结构明显6手动_mo学结构明显

报错堆栈

Epoch 1/15
---------------------------------------------------------------------------
InvalidArgumentError                      Traceback (most recent call last)
~\AppData\Local\Temp\ipykernel_6600\2019445536.py in <module>
     10                         callbacks=[checkpoint_callback, early_stop],
     11                         verbose=1,
---&gt; 12                         class_weight=m_class_weights
     13 #                         class_weight=m_class_weights_np
     14                        )

~\AppData\Roaming\Python\Python37\site-packages\tensorflow_core\python\keras\engine\training.py in fit(self, x, y, batch_size, epochs, verbose, callbacks, validation_split, validation_data, shuffle, class_weight, sample_weight, initial_epoch, steps_per_epoch, validation_steps, validation_freq, max_queue_size, workers, use_multiprocessing, **kwargs)
    725         max_queue_size=max_queue_size,
    726         workers=workers,
--&gt; 727         use_multiprocessing=use_multiprocessing)
    728 
    729   def evaluate(self,

~\AppData\Roaming\Python\Python37\site-packages\tensorflow_core\python\keras\engine\training_generator.py in fit(self, model, x, y, batch_size, epochs, verbose, callbacks, validation_split, validation_data, shuffle, class_weight, sample_weight, initial_epoch, steps_per_epoch, validation_steps, validation_freq, max_queue_size, workers, use_multiprocessing)
    601         shuffle=shuffle,
    602         initial_epoch=initial_epoch,
--&gt; 603         steps_name='steps_per_epoch')
    604 
    605   def evaluate(self,

~\AppData\Roaming\Python\Python37\site-packages\tensorflow_core\python\keras\engine\training_generator.py in model_iteration(model, data, steps_per_epoch, epochs, verbose, callbacks, validation_data, validation_steps, validation_freq, class_weight, max_queue_size, workers, use_multiprocessing, shuffle, initial_epoch, mode, batch_size, steps_name, **kwargs)
    263 
    264       is_deferred = not model._is_compiled
--&gt; 265       batch_outs = batch_function(*batch_data)
    266       if not isinstance(batch_outs, list):
    267         batch_outs = [batch_outs]

~\AppData\Roaming\Python\Python37\site-packages\tensorflow_core\python\keras\engine\training.py in train_on_batch(self, x, y, sample_weight, class_weight, reset_metrics)
   1015       self        Ext MagrightRem提高以至于(画 English)、		mag_right>
   1016       self._make_train_function()
-&gt; 1017       outputs = self.train_function(ins)  # pylint: disable=not-callable
   1018 
   1019     if reset_metrics:

~\AppData\Roaming\Python\Python37\site-packages\tensorflow_core\python\keras\backend.py in __call__(self, inputs)
   3474 
   3475     fetched = self._callable_fn(*array_vals,
-&gt; 3476                                 run_metadata=self.run_metadata)
   3477     self._call_fetch_callbacks(fetched[-len(self._fetches):])
   3478     output_structure = nest.pack_sequence_as(

~\AppData\Roaming\Python\Python37\site-packages\tensorflow_core\python\client\session.py in __call__(self, *args, **kwargs)
   1470         ret = tf_session.TF_SessionRunCallable(self._session._session,
   1471                                                self._handle, args,
-&gt; 1472                                                run_metadata_ptr)
   1473         if run_metadata:
   1474           proto_data = tf_session.TF_GetBuffer(run_metadata_ptr)

InvalidArgumentError: 2 root error(s) found.
  (0) Invalid argument: Incompatible shapes: [32] vs. [32,528]
     [[{{node metrics/acc/Equal}}]]
     [[loss_3/dense_4_loss/weighted_loss/broadcast_weights/assert_broadcastable/is_valid_shape/has_valid_nonscalar_shape/has_invalid_dims/concat/_7061]]
  (1) Invalid argument: Incompatible shapes: [32] vs. [32,528]
     [[{{node metrics/acc/Equal}}]]
0 successful operations.
0 derived errors ignored.

问题分析

报错出现在准确率计算节点metrics/acc/Equal,核心矛盾是模型输出与标签形状不匹配:

  1. [32]是单批次标签的形状(batch size为32),说明标签是整数类别索引;
  2. [32,528]是模型输出经过argmax后的形状,根源是模型结构中多余的Reshape和AveragePooling1D层导致输出变成3D张量:
    • GlobalAveragePooling2D输出是2D张量(None, 1056),但Reshape(None, None, %)上完美磁盘,RTSS的如统计学结构明显,将其拆分为3D张量(None, 528, 1)(1056=2*528);
    • 后续Dense层保持3D维度,最终模型输出为(32,528,7),计算准确率时对最后一维取argmax得到(32,528),与标签的(32,)无法匹配。

解决方案

  1. 修正模型结构:删除多余的Reshape、AveragePooling1D层,全局池化后直接处理2D特征向量:
    from tensorflow.keras import Sequential
    from tensorflow.keras.layers import GlobalAveragePooling2D, BatchNormalization, Dropout, Dense
    
    model = Sequential([
        NASNet(input_shape=(...), include_top=False),
        GlobalAveragePooling2D(),
        BatchNormalization(),
        Dropout(0.5),
        Dense(128, activation='relu'),
        Dropout(0.5),
        Dense(7, activation='softmax')
    ])
    
  2. 匹配损失函数与标签格式:
    • 如果标签是整数索引,改用SparseCategoricalCrossentropy:
      loss = tensorflow.keras.losses.SparseCategoricalCrossentropy(from_logits=False)
      
    • 如果标签是one-hot编码,确保标签形状为(batch_size,7),与模型输出的(batch_size,7)匹配;
  3. 避免动态维度问题:tensorflow-directml在TF1.15对动态维度支持有限,不要用None作为Reshape的参数,如需调整维度请指定固定值。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.09 03:19:52