创建使用PyTorch DataLoader时出现IndexError索引越界报错如何解决
问题排查与解决
报错原因
- 你在
np.loadtxt中指定了分隔符为制表符\t,但你提供的样例数据集实际是通过空格分隔,并非制表符。这导致loadtxt将每一行完整识别为1个字段,你请求读取第0~6列、第7列时就会触发索引越界。
解决方法
可二选一:
方法1(更简便):修改loadtxt的分隔符参数
删除两处np.loadtxt里的delimiter="\t"参数,或者将其改为delimiter=None。np.loadtxt默认会自动匹配任意数量的空白字符(同时兼容空格、制表符分隔的场景)。
修改后的__init__代码如下:
def __init__(self, src_file, num_rows=None): x_tmp = np.loadtxt(src_file, max_rows=num_rows, usecols=range(0,7), skiprows=0, dtype=np.float32) y_tmp = np.loadtxt(src_file, max_rows=num_rows, usecols=7, skiprows=0, dtype=np.long) self.x_data = T.tensor(x_tmp, dtype=T.float32).to(device) self.y_data = T.tensor(y_tmp, dtype=T.long).to(device)
方法2:统一数据分隔符
如果你确实需要使用制表符作为分隔符,将people_train.txt里所有字段之间的分隔符替换为制表符即可。
可选优化建议
你可以一次读取全部数据再拆分,避免重复读取文件:
def __init__(self, src_file, num_rows=None): all_data = np.loadtxt(src_file, max_rows=num_rows, dtype=np.float32) self.x_data = T.tensor(all_data[:, :7], dtype=T.float32).to(device) self.y_data = T.tensor(all_data[:, 7], dtype=T.long).to(device)
内容的提问来源于stack exchange,提问作者Dinesh
相关产品推荐
相关产品推荐

