如何解析包含Tensor字符串列的CSV文件并还原为Tensor?
问题:CSV中Tensor字符串无法转为可用的numpy数组/Tensor
为节省Notebook内存,将图像处理后的数据集保存为CSV文件,但加载后Tensor始终以字符串形式存在,尝试解析时出现以下错误:
ValueError : could not convert string to float: '[-0.4226'
CSV中Tensor字符串示例:
'tensor([[[-0.8849, -0.8849, -0.9192, ..., -1.4329, -2.1179, -2.1179],\n [-0.9192, -0.8849, -0.8678, ..., -1.3987, -2.1179, -2.1179],\n [-0.9020, -0.8849, -0.8678, ..., -1.3644, -2.1179, -2.1179],\n ...,\n [-2.1179, -2.1179, -1.4158, ..., -0.9877, -1.0048, -0.9877],\n [-2.1179, -2.1179, -1.3644, ..., -0.9877, -1.0048, -1.0048],\n [-2.1179, -2.1179, -1.3130, ..., -0.9705, -1.0390, -1.0390]],\n\n [[-0.3200, -0.3200, -0.3550, ..., -1.1779, -2.0357, -2.0357],\n [-0.3550, -0.3200, -0.3025, ..., -1.1604, -2.0357, -2.0357],\n [-0.3375, -0.3200, -0.3200, ..., -1.1604, -2.0357, -2.0357],\n ...,\n [-2.0357, -2.0357, -1.3004, ..., -1.0903, -1.1078, -1.0903],\n [-2.0357, -2.0357, -1.2829, ..., -1.0903, -1.1078, -1.1253],\n [-2.0357, -2.0357, -1.2654, ..., -1.1078, -1.1779, -1.1779]],\n\n [[-1.1944, -1.1770, -1.1770, ..., -1.6650, -1.8044, -1.8044],\n [-1.1944, -1.1596, -1.1247, ..., -1.6476, -1.8044, -1.8044],\n [-1.1421, -1.1247, -1.1073, ..., -1.6302, -1.8044, -1.8044],\n ...,\n [-1.8044, -1.8044, -1.5604, ..., -1.0724, -1.0898, -1.0724],\n [-1.8044, -1.8044, -1.5081, ..., -1.0724, -1.1073, -1.1073],\n [-1.8044, -1.8044, -1.4559, ..., -1.0724, -1.1770, -1.1770]]])'
原尝试的解析代码:
def parse_tensor_string(tensor_string): # Extract the flattened tensor data from the string start_index = tensor_string.find("[[") end_index = tensor_string.find("]]") data_string = tensor_string[start_index+2:end_index] # Remove any extraneous characters data_string = data_string.replace("\n", "").replace(" ", "").replace(",", "") # Convert each substring to a float tensor_list = data_string.split() tensor_floats = [float(s) for s in tensor_list] # Reshape the resulting array to match the original tensor shape tensor_array = np.array(tensor_floats).reshape((3, 224, 224)) # Return the tensor as a numpy array return tensor_array train_set["0"] = train_set["0"].apply(parse_tensor_string)
错误原因
原代码未处理字符串中的[、]和...符号,导致分割后出现[-0.4226这类包含非数字的字符串,无法转为float。
解决方案
方法1:修复原解析函数,处理所有特殊符号
直接在预处理阶段移除所有干扰符号,确保每个分割后的元素都是纯数字:
import numpy as np def parse_tensor_string(tensor_string): # 去掉开头的'tensor('和结尾的')' cleaned_str = tensor_string.replace("tensor(", "").rstrip(")") # 移除所有换行、多余空格、方括号、省略号 cleaned_str = cleaned_str.replace("\n", "").replace(" ", "").replace("[", "").replace("]", "").replace("...", "") # 按逗号分割并过滤空字符串 num_strings = [s.strip() for s in cleaned_str.split(",") if s.strip()] # 转为float数组并reshape tensor_floats = [float(s) for s in num_strings] return np.array(tensor_floats).reshape((3, 224, 224)) train_set["0"] = train_set["0"].apply(parse_tensor_string)
方法2:使用ast.literal_eval安全解析嵌套结构
利用Python的抽象语法树模块,直接将字符串转为嵌套列表,避免手动处理符号:
import ast import numpy as np def parse_tensor_string(tensor_string): # 提取tensor内部的数组字符串 array_str = tensor_string.split("tensor(")[1].rstrip(")") # 用ast解析为嵌套列表 nested_list = ast.literal_eval(array_str) # 转为numpy数组或Tensor return np.array(nested_list) # 如果需要PyTorch Tensor,替换为:return torch.tensor(nested_list) train_set["0"] = train_set["0"].apply(parse_tensor_string)
这种方法更可靠,能自动处理嵌套结构,无需手动处理各种符号。
方法3:优化后续保存方式(避免重复踩坑)
后续保存数据集时,尽量避免存储Tensor的字符串表示:
- 将Tensor转为numpy数组后,扁平化保存为CSV:
np.savetxt("data.csv", tensor.numpy().flatten(), delimiter=","),加载时再reshape:np.loadtxt("data.csv").reshape((3,224,224)) - 使用更适合存储数值数据的格式,比如HDF5(
h5py库)、Pickle,或者PyTorch专用的.pt/.pth格式,这些格式无需手动解析字符串,直接加载即可用。
内容的提问来源于stack exchange,提问作者Michael Macharia
相关产品推荐
相关产品推荐

