使用pandas处理csv并转为张量
基本读取
需要先通过命令安装 pandas 包:
python
pip install pandas
然后通过以下方式,向文件中写入简单的 csv 格式数据:
python
import os
os.makedirs(os.path.join(".","data"),exist_ok=True)
data_file=os.path.join(".","data","stock.csv")
with open(data_file,"w") as f:
f.write('name,price,date\n')
f.write('aaa,12.4,0813\n')
f.write('bbb,NaN,0812\n')
f.write('NaN,NaN,0811\n')
f.write('NaN,31,0810\n')
通过 read_csv 函数读取:
python
import os
import pandas as pd
os.makedirs(os.path.join(".","data"),exist_ok=True)
data_file=os.path.join(".","data","stock.csv")
data=pd.read_csv(data_file)
print(data)
# name price date
# 0 aaa 12.4 813
# 1 bbb NaN 812
# 2 NaN NaN 811
# 3 NaN 31.0 810
数据处理
price 列中包含 NaN 值,我们可以将其填充为该列的平均值:
python
data_me=data.fillna(data.mean(numeric_only=True))
# (12.4+31.0)/2=21.7
print(data_me)
# name price date
# 0 aaa 12.4 813
# 1 bbb 21.7 812
# 2 NaN 21.7 811
# 3 NaN 31.0 810
name 列包含 aaa 、 bbb 、 NaN 三种值,我们可以将其每种值拆成单独一列,并赋一个值:
python
data_me=pd.get_dummies(data_me,dummy_na=True)
print(data_me)
# price date name_aaa name_bbb name_nan
# 0 12.4 813 True False False
# 1 21.7 812 False True False
# 2 21.7 811 False False True
# 3 31.0 810 False False True
然后我们便可以通过 torch 包,来将这个数据转换成多维数组:
python
import torch
x=torch.tensor(data_me.to_numpy(dtype=float))
print(x)
# tensor([[ 12.4000, 813.0000, 1.0000, 0.0000, 0.0000],
# [ 21.7000, 812.0000, 0.0000, 1.0000, 0.0000],
# [ 21.7000, 811.0000, 0.0000, 0.0000, 1.0000],
# [ 31.0000, 810.0000, 0.0000, 0.0000, 1.0000]],
# dtype=torch.float64)