如何在Python中通过列名列表获取列索引并读取CSV生成指定数组?
解决CSV读取与特征、标签分离问题
你的代码存在几个关键问题,我会逐一修正并给出可运行的实现:
代码中的问题点
csv.reader(filename)传参错误:应该传入打开的文件对象而非文件名字符串- 直接遍历
file会读取整行字符串,没有自动分割逗号分隔的字段,应使用csv.reader的迭代器 - 硬编码行数
395不灵活,应该根据实际数据行数动态生成数组 - 未处理表头,无法映射
x_col_names到对应列索引 - 未将字符串类型的字段转换为数值类型(numpy数组需要数值数据)
正确实现代码
假设你已知期末成绩的列名,函数需要接收x_col_names和期末成绩列名作为参数:
import csv import numpy as np def load_student_data(filename, x_col_names, y_col_name): # 存储特征和标签的列表 X_list = [] y_list = [] with open(filename, 'r') as file: reader = csv.reader(file) # 读取表头行 header = next(reader) # 获取特征列的索引 x_indices = [header.index(col) for col in x_col_names] # 获取标签列的索引 y_index = header.index(y_col_name) # 遍历每一行数据 for row in reader: # 提取特征列并转换为数值 x_row = [float(row[idx]) for idx in x_indices] X_list.append(x_row) # 提取标签列并转换为数值 y_val = float(row[y_index]) y_list.append(y_val) # 转换为numpy数组 X = np.array(X_list) y = np.array(y_list) return X, y
使用示例
假设你的列名列表x_col_names包含13个特征列名,期末成绩列名为'G3',调用方式如下:
x_col_names = ['列名1', '列名2', ..., '列名13'] # 替换为你的实际13个特征列名 y_col_name = 'G3' # 替换为你的期末成绩列名 X, y = load_student_data('student.csv', x_col_names, y_col_name) # 验证输出 print(X[0]) # 输出第一行特征 print(X[1]) # 输出第二行特征 print(X.shape, y.shape) # 输出形状 (395,13) (395,)
关键细节说明
- 表头处理:通过
next(reader)获取表头,利用header.index(col)找到目标列的索引,无需硬编码位置 - 数据类型转换:将CSV读取的字符串字段转换为
float(如果是整数可以用int),确保numpy数组为数值类型 - 动态数组:用列表收集数据后再转numpy数组,避免硬编码行数,适配不同数据量
- 索引映射:通过列名索引确保即使CSV列顺序变化,也能正确提取目标字段
内容的提问来源于stack exchange,提问作者user19825372
相关产品推荐
相关产品推荐

