Skip to content

Commit b022a4a

Browse files
esadler-hbowild-endeavor
authored andcommitted
Add pyspark pipeline model transformer (#1101)
* add pyspark pipeline model transformer * reset global spark session in `test_spark_task` * use pytest fixture to reset spark session
1 parent 261b980 commit b022a4a

4 files changed

Lines changed: 98 additions & 1 deletion

File tree

plugins/flytekit-spark/flytekitplugins/spark/__init__.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -17,6 +17,7 @@
1717

1818
from flytekit.configuration import internal as _internal
1919

20+
from .pyspark_transformers import PySparkPipelineModelTransformer
2021
from .schema import SparkDataFrameSchemaReader, SparkDataFrameSchemaWriter, SparkDataFrameTransformer # noqa
2122
from .sd_transformers import ParquetToSparkDecodingHandler, SparkToParquetEncodingHandler
2223
from .task import Spark, new_spark_session # noqa
Lines changed: 45 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,45 @@
1+
import pathlib
2+
from typing import Type
3+
4+
from pyspark.ml import PipelineModel
5+
6+
from flytekit import Blob, BlobMetadata, BlobType, FlyteContext, Literal, LiteralType, Scalar
7+
from flytekit.core.type_engine import TypeEngine
8+
from flytekit.extend import TypeTransformer
9+
10+
11+
class PySparkPipelineModelTransformer(TypeTransformer[PipelineModel]):
12+
_TYPE_INFO = BlobType(format="binary", dimensionality=BlobType.BlobDimensionality.MULTIPART)
13+
14+
def __init__(self):
15+
super(PySparkPipelineModelTransformer, self).__init__(name="PySparkPipelineModel", t=PipelineModel)
16+
17+
def get_literal_type(self, t: Type[PipelineModel]) -> LiteralType:
18+
return LiteralType(blob=self._TYPE_INFO)
19+
20+
def to_literal(
21+
self,
22+
ctx: FlyteContext,
23+
python_val: PipelineModel,
24+
python_type: Type[PipelineModel],
25+
expected: LiteralType,
26+
) -> Literal:
27+
local_path = ctx.file_access.get_random_local_path()
28+
pathlib.Path(local_path).parent.mkdir(parents=True, exist_ok=True)
29+
python_val.save(local_path)
30+
31+
remote_dir = ctx.file_access.get_random_remote_directory()
32+
ctx.file_access.upload_directory(local_path, remote_dir)
33+
34+
return Literal(scalar=Scalar(blob=Blob(uri=remote_dir, metadata=BlobMetadata(type=self._TYPE_INFO))))
35+
36+
def to_python_value(
37+
self, ctx: FlyteContext, lv: Literal, expected_python_type: Type[PipelineModel]
38+
) -> PipelineModel:
39+
local_dir = ctx.file_access.get_random_local_directory()
40+
ctx.file_access.download_directory(lv.scalar.blob.uri, local_dir)
41+
42+
return PipelineModel.load(local_dir)
43+
44+
45+
TypeEngine.register(PySparkPipelineModelTransformer())
Lines changed: 43 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,43 @@
1+
import pyspark
2+
from flytekitplugins.spark import PySparkPipelineModelTransformer
3+
from flytekitplugins.spark.task import Spark
4+
from pyspark.ml import Pipeline, PipelineModel
5+
from pyspark.ml.feature import Imputer
6+
7+
import flytekit
8+
from flytekit import task, workflow
9+
from flytekit.core.type_engine import TypeEngine
10+
11+
12+
def test_type_resolution():
13+
assert type(TypeEngine.get_transformer(PipelineModel)) == PySparkPipelineModelTransformer
14+
15+
16+
def test_pipeline_model_compatibility():
17+
@task(task_config=Spark())
18+
def my_dataset() -> pyspark.sql.DataFrame:
19+
session = flytekit.current_context().spark_session
20+
df = session.createDataFrame([("Megan", 2.0), ("Wayne", float("nan")), ("Dennis", 8.0)], ["name", "age"])
21+
return df
22+
23+
@task(task_config=Spark())
24+
def my_pipleline(df: pyspark.sql.DataFrame) -> PipelineModel:
25+
imputer = Imputer(inputCols=["age"], outputCols=["imputed_age"])
26+
pipeline = Pipeline(stages=[imputer]).fit(df)
27+
return pipeline
28+
29+
@task(task_config=Spark())
30+
def run_pipeline(df: pyspark.sql.DataFrame, pipeline: PipelineModel) -> int:
31+
imputed_df = pipeline.transform(df)
32+
33+
return imputed_df.filter(imputed_df["imputed_age"].isNull()).count()
34+
35+
@workflow
36+
def my_wf() -> int:
37+
df = my_dataset()
38+
pipeline = my_pipleline(df=df)
39+
40+
return run_pipeline(df=df, pipeline=pipeline)
41+
42+
res = my_wf()
43+
assert res == 0

plugins/flytekit-spark/tests/test_spark_task.py

Lines changed: 9 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,6 @@
11
import pandas as pd
22
import pyspark
3+
import pytest
34
from flytekitplugins.spark import Spark
45
from flytekitplugins.spark.task import new_spark_session
56
from pyspark.sql import SparkSession
@@ -10,7 +11,14 @@
1011
from flytekit.core.context_manager import ExecutionParameters, FlyteContextManager
1112

1213

13-
def test_spark_task():
14+
@pytest.fixture(scope="function")
15+
def reset_spark_session() -> None:
16+
pyspark.sql.SparkSession.builder.getOrCreate().stop()
17+
yield
18+
pyspark.sql.SparkSession.builder.getOrCreate().stop()
19+
20+
21+
def test_spark_task(reset_spark_session):
1422
@task(task_config=Spark(spark_conf={"spark": "1"}))
1523
def my_spark(a: str) -> int:
1624
session = flytekit.current_context().spark_session

0 commit comments

Comments
 (0)