<#3648 [BUG] pyspark pipeline type transformer err...
# flyte-github
a
#3648 [BUG] pyspark pipeline type transformer error Issue created by peridotml Describe the bug The pyspark
PipelineModel
transformer works when running spark locally, but is failing to load inside the k8s spark plugin. This is likely because the transformer downloads the pipeline to the driver and it is NOT available on the workers. Likely we need to save and load directly to the remote path. Expected behavior The type transformer shouldn't fail. Additional context to reproduce
Copy code
from flytekit import workflow, task, StructuredDataset
from flytekitplugins.spark import Spark
import flytekit
import pandas as pd
from typing import Tuple

from pyspark.ml.feature import StringIndexer
from pyspark.ml.pipeline import Pipeline, PipelineModel


@task(cache=True, cache_version="1.0")
def create_df() -> pd.DataFrame:
    return pd.DataFrame(
        data={
        "id": [1, 2, 3, 4, 5, 6, 7, 8, 9],
        "cat": ["a", "b", "c", "a", "b", "c", "a", "b", "c"],
        "num": [1, 2, 3, 4, 5, 6, 7, 8, 9],
    })


@task(task_config=Spark(
    spark_conf={
            "spark.driver.memory": "4g",
            "spark.executor.memory": "2g",
            "spark.executor.instances": "1",
            "spark.driver.cores": "2",
            "spark.executor.cores": "1",
    }
), cache=True, cache_version="1.0")
def save_pipeline(df: pd.DataFrame) -> Tuple[pd.DataFrame, PipelineModel]:
    spark = flytekit.current_context().spark_session

    spark_df = spark.createDataFrame(df)

    cat_indexer = StringIndexer(inputCol="cat", outputCol="cat_index")
    pipeline = Pipeline(stages=[cat_indexer])

    fitted_pipeline = pipeline.fit(spark_df)
    spark_df_transformed = fitted_pipeline.transform(spark_df)

    return (
        spark_df_transformed.toPandas(),
        fitted_pipeline
    )

@task(task_config=Spark(
    spark_conf={
        "spark.driver.memory": "4g",
        "spark.executor.memory": "2g",
        "spark.executor.instances": "1",
        "spark.driver.cores": "2",
        "spark.executor.cores": "1",
    }
))
def compare_dfs(pipeline: PipelineModel, df: pd.DataFrame, expected: pd.DataFrame):
    spark = flytekit.current_context().spark_session
    spark_df = spark.createDataFrame(df)
    df_transformed = pipeline.transform(spark_df).toPandas()

    dataframes_equal = expected.sort_values(by=["id"]).equals(
        df_transformed.sort_values(by=["id"]).reset_index(drop=True)
    )

    assert dataframes_equal, "Dataframes are not equal"


@workflow
def test_pipeline_transformer():
    df = create_df()
    df_transformed, pipeline = save_pipeline(df=df)
    compare_dfs(pipeline=pipeline, df=df, expected=df_transformed)
Copy code
[3/3] currentAttempt done. Last Error: SYSTEM::Traceback (most recent call last):

      File "/opt/venv/lib/python3.9/site-packages/flytekit/exceptions/scopes.py", line 165, in system_entry_point
        return wrapped(*args, **kwargs)
      File "/opt/venv/lib/python3.9/site-packages/flytekit/core/base_task.py", line 518, in dispatch_execute
        native_inputs = TypeEngine.literal_map_to_kwargs(exec_ctx, input_literal_map, self.python_interface.inputs)
      File "/opt/venv/lib/python3.9/site-packages/flytekit/core/type_engine.py", line 867, in literal_map_to_kwargs
        return {k: TypeEngine.to_python_value(ctx, lm.literals[k], python_types[k]) for k, v in lm.literals.items()}
      File "/opt/venv/lib/python3.9/site-packages/flytekit/core/type_engine.py", line 867, in <dictcomp>
        return {k: TypeEngine.to_python_value(ctx, lm.literals[k], python_types[k]) for k, v in lm.literals.items()}
      File "/opt/venv/lib/python3.9/site-packages/flytekit/core/type_engine.py", line 831, in to_python_value
        return transformer.to_python_value(ctx, lv, expected_python_type)
      File "/opt/venv/lib/python3.9/site-packages/flytekitplugins/spark/pyspark_transformers.py", line 42, in to_python_value
        return PipelineModel.load(local_dir)
      File "/opt/spark/python/lib/pyspark.zip/pyspark/ml/util.py", line 353, in load
        return cls.read().load(path)
      File "/opt/spark/python/lib/pyspark.zip/pyspark/ml/pipeline.py", line 282, in load
        metadata = DefaultParamsReader.loadMetadata(path, <http://self.sc|self.sc>)
      File "/opt/spark/python/lib/pyspark.zip/pyspark/ml/util.py", line 565, in loadMetadata
        metadataStr = sc.textFile(metadataPath, 1).first()
      File "/opt/spark/python/lib/pyspark.zip/pyspark/rdd.py", line 1906, in first
        raise ValueError("RDD is empty")
Screenshots No response Are you sure this issue hasn't been raised already? ☑︎ Yes Have you read the Code of Conduct? ☑︎ Yes flyteorg/flyte