已将张量转为numpy.ndarray仍报CUDA转numpy错误的技术求助
明明已将张量转为numpy.ndarray,调用plt.scatter仍报CUDA张量转numpy错误?
我已经执行h.detach().clone().cpu().numpy()把张量h转换成了numpy.ndarray,打印type(h)显示结果是<class 'numpy.ndarray'>,而且h的形状是[batch_size,2],但调用plt.scatter时还是抛出了TypeError:can't convert cuda:0 device type tensor to numpy. Use Tensor.cpu() to copy the tensor to host memory first.。错误出现在visualize_embedding函数第7行,相关代码和报错栈如下:
def visualize_embedding(h, color, epoch=None, loss=None): plt.figure(figsize=(7,7)) plt.xticks([]) plt.yticks([]) h = h.detach().clone().cpu().numpy() print(type(h)) plt.scatter(h[:, 0], h[:, 1], s=140, c=color, cmap="Set2") if epoch is not None and loss is not None: plt.xlabel(f'Epoch: {epoch}, Loss: {loss.item():.4f}', fontsize=16) plt.show()
报错信息:
<class 'numpy.ndarray'> --------------------------------------------------------------------------- TypeError Traceback (most recent call last) Cell In[17], line 21 19 loss, h = train(data) 20 if epoch % 10 == 0: ---> 21 visualize_embedding(h, color=data.y, epoch=epoch, loss=loss) 22 time.sleep(0.3) Cell In[16], line 16 14 h = h.detach().clone().cpu().numpy() 15 print(type(h)) ---> 16 plt.scatter(h[:, 0], h[:, 1], s=140, c=color, cmap="Set2") 17 if epoch is not None and loss is not None: 18 plt.xlabel(f'Epoch: {epoch}, Loss: {loss.item():.4f}', fontsize=16) File c:\Users\polyu\Documents\RA\hkjc_dm\hkjc_dm\model\src\venvModel4\lib\site-packages\matplotlib\pyplot.py:3684, in scatter(x, y, s, c, marker, cmap, norm, vmin, vmax, alpha, linewidths, edgecolors, plotnonfinite, data, **kwargs) 3665 @_copy_docstring_and_deprecators(Axes.scatter) 3666 def scatter( 3667 x: float | ArrayLike, (...) 3682 **kwargs, 3683 ) -> PathCollection: -> 3684 __ret = gca().scatter( 3685 x, 3686 y, ... 1030 return self.numpy() 1031 else: -> 1032 return self.numpy().astype(dtype, copy=False) TypeError: can't convert cuda:0 device type tensor to numpy. Use Tensor.cpu() to copy the tensor to host memory first.
问题原因
错误不是来自h,而是来自color参数!你传入的color=data.y是一个CUDA张量,matplotlib在处理c参数时会尝试将其转换为numpy数组,但它还在GPU上,所以触发了这个错误。
解决方案
有两种修复方式:
- 调用函数时直接转换
color:
visualize_embedding(h, color=data.y.detach().cpu().numpy(), epoch=epoch, loss=loss)
- 在
visualize_embedding函数内部处理color,增加兼容性:
在plt.scatter代码前添加:
# 检查是否为PyTorch张量,若是则转成CPU上的numpy数组 if hasattr(color, 'device'): color = color.detach().cpu().numpy()
这样就能解决这个错误了。
内容的提问来源于stack exchange,提问作者Johnny C.
相关产品推荐
相关产品推荐

