【发布时间】:2022-08-11 02:30:16
【问题描述】:
我在 sklearn.pipeline.Pipeline 中有一个训练有素的 Scikit-learn LogisticRegression 模型。这是一项 NLP 任务。模型保存为 pkl 文件(实际上在 ML Studio 模型中,但我将其下载到 databricks dbfs)。
我有一个包含大约 100 万行的 Hive 表(由增量支持)。除其他外,这些行具有ID, 一个关键字上下文列(包含文本),一个建模列(布尔值,表示模型已在该行上运行),以及预言列,它是逻辑回归输出的类的整数。
我的问题是如何更新预测列。
我可以在本地运行
def generatePredictions(data:pd.DataFrame, model:Pipeline) -> pd.DataFrame:
data.loc[:, \'keyword_context\'] = data.keyword_context.apply(lambda x: x.replace(\"\\n\", \" \")
data[\'prediction\'] = model.predict(data.keyword_context)
data[\'modelled\'] = True
return data
这实际上运行得足够快(约 20 秒),但是通过 databricks.sql.connector 将更新运行回数据块需要很多小时。所以我想在 pyspark 笔记本中做同样的事情来绕过冗长的上传。
问题是通常建议使用内置函数(这不是),或者如果必须有一个 udf 那么示例都使用内置类型,而不是管道。我想知道是否应该在函数中加载模型,并且我假设函数需要一行,这意味着很多加载。我真的不确定如何编写函数代码,或者调用它。
标签: pyspark scikit-learn