绘制(1600,1600,1600)三维数据集遇内存不足错误求解决
问题背景
尝试运行以下代码绘制形状为(1600,1600,1600)的三维数据集:
import matplotlib.pyplot as plt from mpl_toolkits.mplot3d import Axes3D import numpy as np # Load the 3D data file data = np.genfromtxt("Ta_parameterspace_2mm.txt", skip_header=14, delimiter=" ", dtype = float) print(data) reflect = data[:,0] emiss = data[:,1] temp = data[:,2] tempdiff = data[:,4] # Create a meshgrid for the x, y, and z coordinates x1, y1, z1 = np.meshgrid(reflect, emiss, temp) # Plot the data using imshow fig = plt.figure() ax = fig.add_subplot(111, projection='3d') ax.imshow(data, extent=[x1.min(), x1.max(), y1.min(), y1.max()], origin='lower', cmap='viridis') plt.show()
执行时触发内存错误:
x1, y1, z1 = np.meshgrid(reflect, emiss, temp) File "<__array_function__ internals>", line 200, in meshgrid File "C:\Users\t.smart\AppData\Local\Packages\PythonSoftwareFoundation.Python.3.10_qbz5n2kfra8p0\LocalCache\local-packages\Python310\site-packages\numpy\lib\function_base.py", line 5045, in meshgrid output = [x.copy() for x in output] File "C:\Users\t.smart\AppData\Local\Packages\PythonSoftwareFoundation.Python.3.10_qbz5n2kfra8p0\LocalCache\local-packages\Python310\site-packages\numpy\lib\function_base.py", line 5045, in <listcomp> output = [x.copy() for x in output] numpy.core._exceptions._ArrayMemoryError: Unable to allocate 30.5 GiB for an array with shape (1600, 1600, 1600) and data type float64
设备仅8GB内存且无管理员权限,需解决该问题。
解决方案
1. 移除不必要的全量meshgrid生成
你代码里的meshgrid完全是冗余操作,生成的(1600,1600,1600)数组直接耗尽内存。改用原始数据直接绘制散点图即可:
import matplotlib.pyplot as plt from mpl_toolkits.mplot3d import Axes3D import numpy as np # 加载数据 data = np.genfromtxt("Ta_parameterspace_2mm.txt", skip_header=14, delimiter=" ", dtype=float) reflect = data[:,0] emiss = data[:,1] tempdiff = data[:,4] fig = plt.figure() ax = fig.add_subplot(111, projection='3d') # 用散点图映射tempdiff的数值 ax.scatter(reflect, emiss, tempdiff, c=tempdiff, cmap='viridis', s=1) # s控制点大小,按需调整 plt.show()
如果数据量仍过大,可抽样减少绘制的点数:
# 每10个点取1个,按需调整抽样间隔 sample_idx = np.arange(0, len(data), 10) reflect_sample = reflect[sample_idx] emiss_sample = emiss[sample_idx] tempdiff_sample = tempdiff[sample_idx] ax.scatter(reflect_sample, emiss_sample, tempdiff_sample, c=tempdiff_sample, cmap='viridis', s=1)
2. 降低数据类型精度
将默认的float64改为float32,内存占用直接减半;精度要求不高时甚至可用float16:
data = np.genfromtxt("Ta_parameterspace_2mm.txt", skip_header=14, delimiter=" ", dtype=np.float32)
3. 分块读取处理数据
如果数据文件极大,避免一次性加载全部数据,分块读取并绘制:
fig = plt.figure() ax = fig.add_subplot(111, projection='3d') # 每次读取1000行,按需调整块大小 for chunk in np.genfromtxt("Ta_parameterspace_2mm.txt", skip_header=14, delimiter=" ", dtype=np.float32, chunksize=1000): reflect_chunk = chunk[:,0] emiss_chunk = chunk[:,1] tempdiff_chunk = chunk[:,4] ax.scatter(reflect_chunk, emiss_chunk, tempdiff_chunk, c=tempdiff_chunk, cmap='viridis', s=1) plt.show()
4. 改用2D热力图替代3D绘图
如果3D视图非必需,将reflect和emiss作为XY轴,tempdiff作为颜色绘制2D热力图,内存占用大幅降低:
import matplotlib.pyplot as plt import numpy as np data = np.genfromtxt("Ta_parameterspace_2mm.txt", skip_header=14, delimiter=" ", dtype=np.float32) reflect = data[:,0] emiss = data[:,1] tempdiff = data[:,4] # 假设数据是规则网格采样,转换为2D网格 x_unique = np.unique(reflect) y_unique = np.unique(emiss) X, Y = np.meshgrid(x_unique, y_unique) Z = tempdiff.reshape(len(y_unique), len(x_unique)) plt.imshow(Z, extent=[x_unique.min(), x_unique.max(), y_unique.min(), y_unique.max()], origin='lower', cmap='viridis') plt.xlabel('reflect') plt.ylabel('emiss') plt.colorbar(label='tempdiff') plt.show()
内容的提问来源于stack exchange,提问作者tjsmert44
相关产品推荐
相关产品推荐

