从CSV读取numpy.ndarray时类型为字符串的解决方法及最优保存方案
问题分析与解决方案
你遇到的核心问题是numpy.ndarray直接存入CSV会被自动转为字符串,读取后自然没法直接用numpy的函数处理。下面分两种情况给你解决办法:
一、修复现有CSV文件的读取问题
如果已经生成了CSV文件,不想重新生成,可以把读取到的字符串格式的数组转换回numpy.ndarray。具体步骤是:
- 把字符串形式的数组(比如
[[1,2],[3,4]]这种格式)用ast.literal_eval()解析成Python列表 - 再把列表转成numpy.ndarray
修改后的读取代码如下:
import ast import numpy as np import matplotlib.pyplot as plt import csv with open('background.csv','r') as csvfile: reader = csv.DictReader(csvfile) for number, value in enumerate(reader): print(number) # 解析字符串为列表,再转成numpy数组 current_img = np.array(ast.literal_eval(value['image'])) print(type(current_img)) # 现在类型是numpy.ndarray plt.imshow(current_img) plt.show()
注意:这种方法只适用于数组的字符串格式是标准Python列表格式的情况,如果你的数组在CSV里存储的格式比较特殊(比如包含array()字样),可能需要先做字符串替换,比如value['image'].replace('array(', '').replace(')', '')再用ast.literal_eval()。
二、更优的numpy数组保存方法(推荐)
CSV其实并不适合存储numpy数组这种结构化数据,推荐用以下几种更高效、更省心的方法:
1. 使用numpy自带的save()和load()
这是最直接的方法,专门用于numpy数组的存储:
import numpy as np # 保存数组 np.save('X_array.npy', X) np.save('Y_array.npy', Y) # 读取数组 X_loaded = np.load('X_array.npy') Y_loaded = np.load('Y_array.npy')
如果要把多个数组存在一个文件里,可以用np.savez():
# 保存多个数组到一个文件 np.savez('background_data.npz', X=X, Y=Y) # 读取 data = np.load('background_data.npz') X_loaded = data['X'] Y_loaded = data['Y']
2. 使用pickle序列化
pickle是Python自带的序列化工具,可以保存几乎所有Python对象,包括numpy数组:
import pickle # 保存 with open('background_data.pkl', 'wb') as f: pickle.dump({'image': X, 'answer': Y}, f) # 读取 with open('background_data.pkl', 'rb') as f: data = pickle.load(f) current_img = data['image'] # 直接是numpy.ndarray类型
3. 如果一定要用CSV(不推荐)
如果因为某些原因必须用CSV,建议把numpy数组扁平化后存储,读取时再恢复形状:
import csv import numpy as np import ast # 保存时扁平化数组,同时记录原形状 X_flat = X.flatten() with open('background.csv','w',newline='') as csvfile: fieldnames = ['image_shape', 'image_flat', 'answer'] writer = csv.DictWriter(csvfile, fieldnames=fieldnames) writer.writeheader() writer.writerow({ 'image_shape': str(X.shape), 'image_flat': X_flat.tolist(), 'answer': Y.tolist() }) # 读取时恢复形状 with open('background.csv','r') as csvfile: reader = csv.DictReader(csvfile) for value in reader: img_shape = tuple(ast.literal_eval(value['image_shape'])) current_img = np.array(ast.literal_eval(value['image_flat'])).reshape(img_shape) print(type(current_img)) plt.imshow(current_img) plt.show()
总结一下,优先选择numpy的save()/savez()或者pickle,它们不仅操作简单,而且保存和读取的效率远高于CSV,还能完整保留数组的类型和形状信息,不会出现格式转换的问题。
内容的提问来源于stack exchange,提问作者HJRaBob
相关产品推荐
相关产品推荐

