【问题标题】:How do I unit test PySpark programs?如何对 PySpark 程序进行单元测试?
【发布时间】:2016-02-22 01:36:46
【问题描述】:

我当前的 Java/Spark 单元测试方法通过使用“本地”实例化 SparkContext 并使用 JUnit 运行单元测试来工作(详细信息 here)。

必须组织代码在一个函数中执行 I/O,然后使用多个 RDD 调用另一个函数。

这很好用。我有一个用 Java + Spark 编写的经过高度测试的数据转换。

我可以用 Python 做同样的事情吗?

如何使用 Python 运行 Spark 单元测试?

【问题讨论】:

标签: python unit-testing apache-spark pyspark


【解决方案1】:

我也建议使用 py.test。 py.test 可以轻松创建可重用的 SparkContext 测试夹具并使用它来编写简洁的测试函数。您还可以专门化固定装置(例如创建 StreamingContext)并在测试中使用其中的一个或多个。

我在 Medium 上写了一篇关于这个主题的博文:

https://engblog.nextdoor.com/unit-testing-apache-spark-with-py-test-3b8970dc013b

这是帖子中的一个sn-p:

pytestmark = pytest.mark.usefixtures("spark_context")
def test_do_word_counts(spark_context):
    """ test word couting
    Args:
       spark_context: test fixture SparkContext
    """
    test_input = [
        ' hello spark ',
        ' hello again spark spark'
    ]

    input_rdd = spark_context.parallelize(test_input, 1)
    results = wordcount.do_word_counts(input_rdd)

    expected_results = {'hello':2, 'spark':3, 'again':1}  
    assert results == expected_results

【讨论】:

  • 欢迎来到 SO!主要是链接答案不受欢迎。 (也就是说,如果链接消失,答案将没有持久的价值。)建议添加一些有用的文本来总结或突出链接资源中的关键点。
  • @Vikas Kawadia 你能看看https://stackoverflow.com/questions/49420660/unit-test-pyspark-code-using-python
  • 博文中概述的RDD测试很好,但是DataFrame测试只检查有两行数据。它不会验证 DataFrame 模式和内容是否相同,因此它不是一个健壮的测试。有关进行 DataFrame 比较的更好方法,请参阅我的答案。
【解决方案2】:

如果您使用的是 Spark 2.x 和 SparkSession,这里有一个 pytest 解决方案。我也在导入第三方包。

import logging

import pytest
from pyspark.sql import SparkSession

def quiet_py4j():
    """Suppress spark logging for the test context."""
    logger = logging.getLogger('py4j')
    logger.setLevel(logging.WARN)


@pytest.fixture(scope="session")
def spark_session(request):
    """Fixture for creating a spark context."""

    spark = (SparkSession
             .builder
             .master('local[2]')
             .config('spark.jars.packages', 'com.databricks:spark-avro_2.11:3.0.1')
             .appName('pytest-pyspark-local-testing')
             .enableHiveSupport()
             .getOrCreate())
    request.addfinalizer(lambda: spark.stop())

    quiet_py4j()
    return spark


def test_my_app(spark_session):
   ...

注意,如果使用 Python 3,我必须将其指定为 PYSPARK_PYTHON 环境变量:

import os
import sys

IS_PY2 = sys.version_info < (3,)

if not IS_PY2:
    os.environ['PYSPARK_PYTHON'] = 'python3'

否则会报错:

异常:worker 中的 Python 与 2.7 版本不同 驱动程序 3.5,PySpark 无法使用不同的次要版本运行。请 检查环境变量 PYSPARK_PYTHON 和 PYSPARK_DRIVER_PYTHON 设置正确。

【讨论】:

  • 当我在 Spark 2.0.2 上使用此代码时,avro 插件不起作用
  • Avro 插件可以像使用 Spark 2.1 一样加载,但不能使用 Spark 2.0.2。在您尝试使用 Avro 格式之前,您不会收到错误消息。我自己测试过。
  • 一种更简单的设置 PYSPARK_PYTHON 正确值的方法:os.environ['PYSPARK_PYTHON'] = sys.executable - 这将设置为当前运行的 python 的值,并且希望能更好地处理 venvs
  • @ksindi 你能看看https://stackoverflow.com/questions/49420660/unit-test-pyspark-code-using-python
  • @user9367133 回答了你的问题
【解决方案3】:

假设你已经安装了pyspark,你可以使用下面的类在unittest中对它进行单元测试:

import unittest
import pyspark


class PySparkTestCase(unittest.TestCase):

    @classmethod
    def setUpClass(cls):
        conf = pyspark.SparkConf().setMaster("local[2]").setAppName("testing")
        cls.sc = pyspark.SparkContext(conf=conf)
        cls.spark = pyspark.SQLContext(cls.sc)

    @classmethod
    def tearDownClass(cls):
        cls.sc.stop()

例子:

class SimpleTestCase(PySparkTestCase):

    def test_with_rdd(self):
        test_input = [
            ' hello spark ',
            ' hello again spark spark'
        ]

        input_rdd = self.sc.parallelize(test_input, 1)

        from operator import add

        results = input_rdd.flatMap(lambda x: x.split()).map(lambda x: (x, 1)).reduceByKey(add).collect()
        self.assertEqual(results, [('hello', 2), ('spark', 3), ('again', 1)])

    def test_with_df(self):
        df = self.spark.createDataFrame(data=[[1, 'a'], [2, 'b']], 
                                        schema=['c1', 'c2'])
        self.assertEqual(df.count(), 2)

请注意,这会为每个类创建一个上下文。使用setUp 而不是setUpClass 来获取每个测试的上下文。这通常会在执行测试时增加大量开销时间,因为目前创建新的 Spark 上下文非常昂贵。

【讨论】:

    【解决方案4】:

    我使用pytest,它允许测试夹具,因此您可以实例化 pyspark 上下文并将其注入到所有需要它的测试中。类似于

    @pytest.fixture(scope="session",
                    params=[pytest.mark.spark_local('local'),
                            pytest.mark.spark_yarn('yarn')])
    def spark_context(request):
        if request.param == 'local':
            conf = (SparkConf()
                    .setMaster("local[2]")
                    .setAppName("pytest-pyspark-local-testing")
                    )
        elif request.param == 'yarn':
            conf = (SparkConf()
                    .setMaster("yarn-client")
                    .setAppName("pytest-pyspark-yarn-testing")
                    .set("spark.executor.memory", "1g")
                    .set("spark.executor.instances", 2)
                    )
        request.addfinalizer(lambda: sc.stop())
    
        sc = SparkContext(conf=conf)
        return sc
    
    def my_test_that_requires_sc(spark_context):
        assert spark_context.textFile('/path/to/a/file').count() == 10
    

    然后您可以通过调用py.test -m spark_local 或在YARN 中使用py.test -m spark_yarn 在本地模式下运行测试。这对我来说效果很好。

    【讨论】:

    • 你能看看https://stackoverflow.com/questions/49420660/unit-test-pyspark-code-using-python
    【解决方案5】:

    您可以通过在测试套件中的 DataFrame 上运行您的代码并比较 DataFrame 列相等或两个整个 DataFrame 的相等来测试 PySpark 代码。

    quinn project has several examples

    为测试套件创建 SparkSession

    使用此夹具创建一个 tests/conftest.py 文件,以便您可以在测试中轻松访问 SparkSession。

    import pytest
    from pyspark.sql import SparkSession
    
    @pytest.fixture(scope='session')
    def spark():
        return SparkSession.builder \
          .master("local") \
          .appName("chispa") \
          .getOrCreate()
    

    列相等示例

    假设您想测试以下从字符串中删除所有非单词字符的函数。

    def remove_non_word_characters(col):
        return F.regexp_replace(col, "[^\\w\\s]+", "")
    

    您可以使用chispa 库中定义的assert_column_equality 函数来测试此函数。

    def test_remove_non_word_characters(spark):
        data = [
            ("jo&&se", "jose"),
            ("**li**", "li"),
            ("#::luisa", "luisa"),
            (None, None)
        ]
        df = spark.createDataFrame(data, ["name", "expected_name"])\
            .withColumn("clean_name", remove_non_word_characters(F.col("name")))
        assert_column_equality(df, "clean_name", "expected_name")
    

    DataFrame 相等示例

    有些功能需要通过比较整个 DataFrame 来进行测试。这是一个对 DataFrame 中的列进行排序的函数。

    def sort_columns(df, sort_order):
        sorted_col_names = None
        if sort_order == "asc":
            sorted_col_names = sorted(df.columns)
        elif sort_order == "desc":
            sorted_col_names = sorted(df.columns, reverse=True)
        else:
            raise ValueError("['asc', 'desc'] are the only valid sort orders and you entered a sort order of '{sort_order}'".format(
                sort_order=sort_order
            ))
        return df.select(*sorted_col_names)
    

    这是你要为这个函数编写的一个测试。

    def test_sort_columns_asc(spark):
        source_data = [
            ("jose", "oak", "switch"),
            ("li", "redwood", "xbox"),
            ("luisa", "maple", "ps4"),
        ]
        source_df = spark.createDataFrame(source_data, ["name", "tree", "gaming_system"])
    
        actual_df = T.sort_columns(source_df, "asc")
    
        expected_data = [
            ("switch", "jose", "oak"),
            ("xbox", "li", "redwood"),
            ("ps4", "luisa", "maple"),
        ]
        expected_df = spark.createDataFrame(expected_data, ["gaming_system", "name", "tree"])
    
        assert_df_equality(actual_df, expected_df)
    

    测试 I/O

    通常最好从 I/O 函数中抽象出代码逻辑,这样更容易测试。

    假设你有这样一个函数。

    def your_big_function:
        df = spark.read.parquet("some_directory")
        df2 = df.withColumn(...).transform(function1).transform(function2)
        df2.write.parquet("other directory")
    

    最好像这样重构代码:

    def all_logic(df):
      return df.withColumn(...).transform(function1).transform(function2)
    
    def your_formerly_big_function:
        df = spark.read.parquet("some_directory")
        df2 = df.transform(all_logic)
        df2.write.parquet("other directory")
    

    这样设计您的代码可以让您轻松测试all_logic 函数与上面提到的列相等或DataFrame 相等函数。您可以使用模拟来测试your_formerly_big_function。通常最好在测试套件中避免 I/O(但有时是不可避免的)。

    【讨论】:

    • assert_df_equality 未找到
    【解决方案6】:

    pyspark 有 unittest 模块,可以如下使用

    from pyspark.tests import ReusedPySparkTestCase as PySparkTestCase
    
    class MySparkTests(PySparkTestCase):
        def spark_session(self):
            return pyspark.SQLContext(self.sc)
    
        def createMockDataFrame(self):
             self.spark_session().createDataFrame(
                [
                    ("t1", "t2"),
                    ("t1", "t2"),
                    ("t1", "t2"),
                ],
                ['col1', 'col2']
            )
    

    【讨论】:

      【解决方案7】:

      前段时间我也遇到过同样的问题,在阅读了几篇文章、论坛和一些 StackOverflow 答案后,我最终为 pytest 编写了一个小插件:pytest-spark

      我已经使用它几个月了,一般的工作流程在 Linux 上看起来不错:

      1. 安装 Apache Spark(设置 JVM + 将 Spark 的分发解压到某个目录)
      2. 安装“pytest”+插件“pytest-spark”
      3. 在您的项目目录中创建“pytest.ini”并在其中指定 Spark 位置。
      4. 像往常一样通过 pytest 运行测试。
      5. 您可以选择在测试中使用由插件提供的夹具“spark_context” - 它会尽量减少 Spark 在输出中的日志。

      【讨论】:

        猜你喜欢
        • 1970-01-01
        • 1970-01-01
        • 2016-04-04
        • 1970-01-01
        • 1970-01-01
        • 1970-01-01
        • 2012-01-08
        • 1970-01-01
        • 1970-01-01
        相关资源
        最近更新 更多