使用SparkTrials并行化Hyperopt时,封装为类后出现SparkContext引用错误
我太懂你这种挫败感了——明明用普通函数写的时候跑的好好的,一封装成类就弹出那个RuntimeError: It appears that you are attempting to reference SparkContext from a broadcast variable...的报错,简直摸不着头脑对吧?
问题根源
其实这个报错的核心原因很简单:当你把类的方法self._rmse传给fmin时,整个类实例self会被序列化后传到Spark worker节点。而你的类初始化时保存了self.spark(SparkSession/SparkContext对象),这个对象是Spark架构中严格属于driver端的,根本不能被序列化传到worker,Spark明确禁止在worker代码里直接引用driver的SparkContext,这就是触发报错的直接原因。
而且你会发现,哪怕你在_rmse里根本没用到self.spark也没用——只要类实例里包含这个对象,序列化的时候就会被带过去,进而触发检查报错。
解决办法
根据你的代码场景,我给你几个可行的解决方案:
方案1:移除类中对Spark对象的引用(最适合你的场景)
看你的代码,_load_data用的是deltalake Python库直接读取DBFS路径,完全不需要SparkSession!那类里根本没必要保存self.spark,直接把它从__init__里删掉就行:
class SomeClass: def __init__(self, path): # 去掉spark参数,不再持有Spark实例 self.path = path def model(self, data, a, b): return a + b * data def rmse(self, x, y, a, b): e = self.model(x, a, b) - y return np.sqrt(np.mean(e ** 2)) def _load_data(self): return DeltaTable(f'/dbfs/{self.path}').to_pandas() def _rmse(self, params): data = self._load_data() x = data.loc[:, 'x'].to_numpy() y = data.loc[:, 'y'].to_numpy() return self.rmse(x, y, **params) def run_hyperopt(self): trials = SparkTrials() space = {'a': hp.uniform('a', 0, 1), 'b': hp.uniform('b', -1,1)} best = fmin(fn=self._rmse, space=space, algo=tpe.suggest, max_evals=10, trials=trials) return best # 调用方式也简化了 path = '/FileStore/mv/sample_data/random_x' c = SomeClass(path) c.run_hyperopt()
这样类实例序列化时就不会包含Spark对象,worker端执行_rmse时自然不会触发那个错误。
方案2:提前在driver端加载数据(适合需要用Spark加载数据的场景)
如果之后你需要用Spark加载数据,那一定要把数据加载操作放在driver端完成,再把加载好的数据传给类,不要让worker去碰Spark:
class SomeClass: def __init__(self, x, y): # 直接传入driver端处理好的x、y数组 self.x = x self.y = y def model(self, data, a, b): return a + b * data def rmse(self, x, y, a, b): e = self.model(x, a, b) - y return np.sqrt(np.mean(e ** 2)) def _rmse(self, params): # 直接用提前加载好的数据,完全不涉及Spark return self.rmse(self.x, self.y, **params) def run_hyperopt(self): trials = SparkTrials() space = {'a': hp.uniform('a', 0, 1), 'b': hp.uniform('b', -1,1)} best = fmin(fn=self._rmse, space=space, algo=tpe.suggest, max_evals=10, trials=trials) return best # 在driver端提前加载并处理数据 path = '/FileStore/mv/sample_data/random_x' data = DeltaTable(f'/dbfs/{path}').to_pandas() x = data.loc[:, 'x'].to_numpy() y = data.loc[:, 'y'].to_numpy() c = SomeClass(x, y) c.run_hyperopt()
这种方式下,worker端只需要计算RMSE,完全不接触SparkContext,从根源上避免了报错。
方案3:使用静态方法剥离类实例依赖
如果你的方法不需要访问类的其他属性(除了路径),可以把核心方法改成静态方法,这样就不需要依赖整个类实例,也就不会序列化Spark对象:
class SomeClass: def __init__(self, path): self.path = path @staticmethod def model(data, a, b): return a + b * data @staticmethod def rmse(x, y, a, b): e = SomeClass.model(x, a, b) - y return np.sqrt(np.mean(e ** 2)) @staticmethod def _load_data(path): return DeltaTable(f'/dbfs/{path}').to_pandas() def _rmse(self, params): data = SomeClass._load_data(self.path) x = data.loc[:, 'x'].to_numpy() y = data.loc[:, 'y'].to_numpy() return SomeClass.rmse(x, y, **params) def run_hyperopt(self): trials = SparkTrials() space = {'a': hp.uniform('a', 0, 1), 'b': hp.uniform('b', -1,1)} best = fmin(fn=self._rmse, space=space, algo=tpe.suggest, max_evals=10, trials=trials) return best
这个方案的核心还是避免让类实例携带Spark对象,本质和方案1是一致的。
总结
不管用哪种方案,核心思路都是不让Spark worker端接触到driver的SparkContext对象——要么移除类中对Spark的引用,要么把Spark相关操作都放在driver端完成,别让worker去碰。
备注:内容来源于stack exchange,提问作者deblue

