使用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, ---> 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, --> 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, --> 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 --> 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() -> 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, -> 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, -> 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,核心矛盾是模型输出与标签形状不匹配:
[32]是单批次标签的形状(batch size为32),说明标签是整数类别索引;[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,)无法匹配。
解决方案
- 修正模型结构:删除多余的
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') ]) - 匹配损失函数与标签格式:
- 如果标签是整数索引,改用
SparseCategoricalCrossentropy:loss = tensorflow.keras.losses.SparseCategoricalCrossentropy(from_logits=False) - 如果标签是one-hot编码,确保标签形状为
(batch_size,7),与模型输出的(batch_size,7)匹配;
- 如果标签是整数索引,改用
- 避免动态维度问题:tensorflow-directml在TF1.15对动态维度支持有限,不要用
None作为Reshape的参数,如需调整维度请指定固定值。
内容的提问来源于stack exchange,提问作者user21525821
相关产品推荐
相关产品推荐

