TensorFlow2.1+Python3.7下Numpy数据转GPU张量失败求助
解决TensorFlow 2.1中tf.data.Dataset无法将自定义数据转移到GPU的问题
首先看你的代码和运行日志,核心问题大概率是数据类型不匹配导致TensorFlow无法自动将数据转移到GPU,进而所有运算都卡在CPU执行(GPU利用率为0);而随机生成的数据刚好匹配模型期望的类型,所以能正常调用GPU。下面是具体的分析和修复步骤:
1. 修复数据类型不匹配的核心问题
这是最关键的触发点:
- 在
TermWeightDataset的__init__方法中,你将labels处理为float64类型:self.labels = tf.keras.preprocessing.sequence.pad_sequences(self.labels,maxlen=self.max_count,padding='post',value=0.0,dtype="float64") - 但在
create_datasets方法里,你给from_generator指定的输出类型是(tf.int64, tf.float32),这就导致label的实际类型和Dataset声明的类型不一致。TensorFlow会在CPU上做隐式类型转换,后续的模型运算也会被迫留在CPU执行。
修复方式:
统一label的数据类型为float32(TensorFlow默认更推荐用float32,运算更快且节省GPU内存):
self.labels = tf.keras.preprocessing.sequence.pad_sequences(self.labels,maxlen=self.max_count,padding='post',value=0.0,dtype="float32")
2. 优先使用更高效的Dataset创建方式
你代码里注释掉了from_tensor_slices的实现,其实这种方式比from_generator更适合处理已加载到内存的numpy数据:
from_tensor_slices可以直接将numpy数组转换为TensorFlow张量,TensorFlow能更好地自动处理设备(GPU/CPU)的分配,避免from_generator带来的Python线程开销和设备转移问题。- 替换成以下代码:
def create_datasets(self): return tf.data.Dataset.from_tensor_slices((self.data,self.labels)).batch(self.batch_size,drop_remainder=True)
3. 显式指定设备(可选,自动分配失效时)
如果上述修复后还是无法使用GPU,可以显式强制将数据和模型放到GPU:
方式一:在训练时指定设备
修改train函数中的batch处理部分:
for idx, (inputs_batch,labels_batch) in enumerate(dataset): with tf.device('/GPU:0'): inputs_batch = tf.identity(inputs_batch) labels_batch = tf.identity(labels_batch) logits,loss = train_one_step(model,inputs_batch,labels_batch,loss_function,optimizer)
方式二:将模型移到GPU
初始化模型后,显式转移到GPU:
model = LSTMBasedModel(vocab_size, input_dim, hiddien_dim, output_dim, embedding_matrix) # 先构建模型(传入一个示例输入触发层初始化) _ = model(tf.random.uniform((1, max_count), dtype=tf.int64)) model = model.to('/GPU:0')
4. 完善padded_batch的padding参数(如果继续使用from_generator)
如果你坚持使用from_generator,需要确保padded_batch的padding值类型和数据匹配,避免额外的类型转换:
return tf.data.Dataset.from_generator( self.generate, (tf.int64,tf.float32) ).padded_batch( self.batch_size, padded_shapes=(self.max_count,self.max_count), padding_values=(0, 0.0) # 输入是int类型用0,label是float类型用0.0 )
最后,日志末尾的Killed提示可能是内存不足导致的,可以尝试减小batch_size或者优化数据加载流程,不过这是次要问题,先解决GPU利用率的问题即可。
内容的提问来源于stack exchange,提问作者Peter
相关产品推荐
相关产品推荐

