Python跨类传递for循环迭代变量后如何在类内访问该值
问题根因
你的代码存在3个核心语法/逻辑错误,导致传入cross类的迭代计数无法正常调用:
- 类构造方法拼写错误:Python类的构造方法为固定名称
__init__(前后各2个下划线),你写的init__缺少前置双下划线,实例化类时根本不会触发初始化逻辑,传入的g值没有绑定到类实例上 - 变量作用域错误:你在KFold循环中直接判断
v == c,但v是cross类的实例属性,脱离类实例直接访问v会直接触发变量未定义报错 - 缩进逻辑错误:你把KFold拆分、数据集初始化的逻辑写在了
cross类定义的同级缩进下,不属于类的内部逻辑,自然无法访问类内部存储的传入值
修正方案
首先修正cross类的构造方法拼写,把需要用到传入计数的训练相关逻辑收拢到类的实例方法中,通过self.前缀访问实例存储的计数值。
修正后的cross类定义
class cross(): # 构造方法必须为__init__,前后双下划线 def __init__(self, value): # 传入的迭代计数绑定为实例属性self.v,类内部所有方法都可以通过self.v访问该值 self.v = value print('i entered the classs') # 初始化时自动执行KFold拆分逻辑 self.run_kfold_pipeline() def run_kfold_pipeline(self): fold_count = 0 for train_index, test_index in kf.split(x_trainval): fold_count += 1 x_train, x_val = x[train_index], x[test_index] y_train, y_val = y[train_index], y[test_index] print("TRAIN:", train_index, "TEST:", test_index) print("epoch", fold_count) # 直接访问实例属性self.v,就是外层循环传入的g值 if self.v == fold_count: print("compare") break # 数据集、DataLoader初始化逻辑收拢到类方法内,可正常拿到拆分后的数据 train_dataset = classifierdataset(torch.from_numpy(x_train).float(), torch.from_numpy(y_train).long()) val_dataset = classifierdataset(torch.from_numpy(x_val).float(), torch.from_numpy(y_val).long()) test_dataset = classifierdataset(torch.from_numpy(x_test).float(), torch.from_numpy(y_test).long()) train_loader = DataLoader(dataset=train_dataset, batch_size=BATCH_SIZE) val_loader = DataLoader(dataset = val_dataset, batch_size = 1) test_loader = DataLoader(dataset = test_dataset , batch_size = 1) print("train",x_train,"val",x_val)
外层迭代代码无需调整
你原来写的外层循环生成迭代计数、实例化类的逻辑可以直接运行:
for e in range(1, EPOCHS+1): g = 0 for j in range(jin): g += 1 print(g) train_epoch_loss = 0 train_epoch_acc = 0 # 实例化时传入g值,会自动执行初始化和KFold拆分逻辑 p1 = cross(g) model.train()
注意:Python类中需要跨方法共享的实例属性必须加
self.前缀,直接写变量名会被识别为方法内的局部变量,跨作用域访问就会触发报错。
内容的提问来源于stack exchange,提问作者luffy
相关产品推荐
相关产品推荐

