如何使用assert语句限制输入列号不超过数据集的实际总列数
解决方案
你需要补充的断言语句可直接通过numpy数组的shape属性获取总列数,同时兼容Python正负索引的使用规则,修改后的完整代码如下:
import numpy as np data = np.loadtxt('data.csv', delimiter = ',', skiprows = 1) def get_column(data, col_num): assert len(data.shape) == 2, "输入数据集必须为二维数组" assert type(col_num) is int, "列号参数必须为整数" # 新增列号边界校验 assert -data.shape[1] <= col_num < data.shape[1], f"列号超出合法范围,当前数据集共有{data.shape[1]}列" return data[:, col_num]
校验逻辑说明
- 二维numpy数组的
shape属性返回(行数, 列数),data.shape[1]可直接读取总列数,无需提前知晓数据集的具体结构 - 校验范围同时覆盖正负索引的合法区间:
- 正索引要求大于等于0,小于总列数
- 负索引要求大于等于负的总列数,匹配Python从末尾倒序索引的规则
- 断言内置的报错提示可以快速明确参数问题,方便定位传入的列号和实际列数的差异
另外你原有代码中的第一个assert语句缺少闭合右括号,上述代码已经同步修正了该语法问题。
内容的提问来源于stack exchange,提问作者Bayle
相关产品推荐
相关产品推荐

