【问题标题】:Loading pickled model with Pyspark使用 Pyspark 加载腌制模型
【发布时间】: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


    【解决方案1】:

    为了记录,我的问题只是我没有在 turn_labeller 方法的定义上方添加 @staticmethod 装饰器。一个简单的修复。

    【讨论】:

      猜你喜欢
      • 2018-08-03
      • 2017-03-01
      • 2021-12-24
      • 1970-01-01
      • 2021-06-25
      • 2019-08-17
      • 2018-06-26
      • 1970-01-01
      • 2013-08-18
      相关资源
      最近更新 更多