【发布时间】:2018-04-04 01:48:36
【问题描述】:
我正在尝试使用 PySpark 从 S3 加载腌制模型,然后使用该模型进行预测。我可以很好地加载模型,但是当我尝试将模型提供给进行预测的方法时,我遇到了PicklingError: Cannot pickle files that are not opened for reading 我已经阅读了关于什么可以腌制和不能腌制的文档,但我不能似乎找到了错误。
加载模型的代码:
rdd_pickle = spark.sparkContext.binaryFiles(model_path_in_s3)
l = rdd_pickle.collect()
pickle_text = l[0][1]
self.model = pickle.loads(pickle_text)
模型进行预测的方法:
def turn_labeller(convo):
"""Annotate a conversation with turn labels.
:param convo: Conversation whose turns haven't been labelled.
:type convo: list of dict
:param model: CRF model used to predict turn labels
:type model: sklearn-crfsuite.CRF
:return: The convo, now with labelled turns
:rtype: list of dict
"""
turn_features = [extract_turn_features(i, convo) for i in range(len(convo))]
predicted_labels = model.predict_single(turn_features)
for i,turn in enumerate(convo):
if i == 0:
turn["previous_turn_label"] = "__ROOT__"
else:
turn["previous_turn_label"] = predicted_labels[i-1]
turn["turn_label"] = predicted_labels[i]
return convo
以及进行所有计算的方法:
def run(self):
"""Run the pipeline."""
# Ok, so next thing is to run our transformations
rdd_tagged = (
self.rdd_interactions
.filter(lambda d, valid=self.filter_invalid_thread_id: "Thread ID" in d.keys() and valid(d["Thread ID"]))
.filter(lambda d, config=self.regex_configs: d["Thread ID"] not in config["BadThreads"])
.map(lambda d: (d["Thread ID"], d))
.groupByKey()
.map(lambda t: list(t[1])) # Now all interactions per conversation are together
.filter(lambda ld: len(ld) > 0)
.map(lambda ld, clean_conv=self.clean_conversation: clean_conv(ld))
.filter(lambda d: d is not None)
.filter(lambda d: d["conversation"]) # remove conversations without content
.map(lambda d, segment=self.turn_segmentation, conf=self.regex_configs: segment(d, conf))
.flatMap(lambda ld, label=self.turn_labeller: label(ld))
)
一切都运行到调用 turn_labeller 的最后一个 flatMap。导致错误的那个调用是怎么回事?
【问题讨论】:
标签: scikit-learn pyspark pickle