You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Pandas ExtensionArray极简实现示例:支持NA的整数数组并通过官方测试

在我看来,Pandas ExtensionArray 是非常需要简单入门示例的场景之一,但我一直没有找到足够简洁的参考案例。

创建 ExtensionArray 的要求

如需创建 ExtensionArray,需要完成以下两项操作:

  • 实现 ExtensionDtype 并完成注册
  • 继承 ExtensionArray 并实现所有要求的必要方法

Pandas官方文档的扩展类型章节有相关基础概述。

现有实现参考

目前公开的ExtensionArray实现案例有很多:

  • Pandas内置的内部扩展数组
  • Geopandas 提供的 GeometryArray
  • Pandas官方文档中提供的扩展数据类型项目列表
    • 例如 CyberPandas 的 IPArray
  • 其他开源实现,例如 Fletcher 的 StringSupportingExtensionArray

需求说明

尽管我已经调研了上述所有案例,但依然觉得扩展数组的实现逻辑难以理解:所有案例都包含大量定制化的业务功能,很难剥离出ExtensionArray的基础必要实现部分,我相信很多开发者都遇到了同样的问题。

因此我需要一个极简、可运行的ExtensionArray实现示例,要求该类能够通过Pandas官方提供的所有测试用例,确保ExtensionArray的行为符合规范,下方已附上我基于官方测试规范编写的验证代码。

为了让需求更明确,本次需要实现的是继承ExtensionArray的可空整数数组,本质是精简版的IntegerArray,仅保留ExtensionArray的基础功能,无需额外定制能力。


解决方案测试规范

我使用以下fixture和测试用例验证解决方案的有效性,所有测试逻辑均基于Pandas官方文档要求编写:

import operator

import numpy as np
from pandas import Series
import pytest

from pandas.tests.extension.base.casting import BaseCastingTests  # noqa
from pandas.tests.extension.base.constructors import BaseConstructorsTests  # noqa
from pandas.tests.extension.base.dtype import BaseDtypeTests  # noqa
from pandas.tests.extension.base.getitem import BaseGetitemTests  # noqa
from pandas.tests.extension.base.groupby import BaseGroupbyTests  # noqa
from pandas.tests.extension.base.interface import BaseInterfaceTests  # noqa
from pandas.tests.extension.base.io import BaseParsingTests  # noqa
from pandas.tests.extension.base.methods import BaseMethodsTests  # noqa
from pandas.tests.extension.base.missing import BaseMissingTests  # noqa
from pandas.tests.extension.base.ops import (  # noqa
    BaseArithmeticOpsTests,
    BaseComparisonOpsTests,
    BaseOpsUtil,
    BaseUnaryOpsTests,
)
from pandas.tests.extension.base.printing import BasePrintingTests  # noqa
from pandas.tests.extension.base.reduce import (  # noqa
    BaseBooleanReduceTests,
    BaseNoReduceTests,
    BaseNumericReduceTests,
)
from pandas.tests.extension.base.reshaping import BaseReshapingTests  # noqa
from pandas.tests.extension.base.setitem import BaseSetitemTests  # noqa

from .extension import NullableIntArray


@pytest.fixture
def dtype():
    """A fixture providing the ExtensionDtype to validate."""
    return 'NullableInt'


@pytest.fixture
def data():
    """
    Length-100 array for this type.
    * data[0] and data[1] should both be non missing
    * data[0] and data[1] should not be equal
    """
    return NullableIntArray(np.array(list(range(100))))


@pytest.fixture
def data_for_twos():
    """Length-100 array in which all the elements are two."""
    return NullableIntArray(np.array([2] * 2))


@pytest.fixture
def data_missing():
    """Length-2 array with [NA, Valid]"""
    return NullableIntArray(np.array([np.nan, 2]))


@pytest.fixture(params=["data", "data_missing"])
def all_data(request, data, data_missing):
    """Parametrized fixture giving 'data' and 'data_missing'"""
    if request.param == "data":
        return data
    elif request.param == "data_missing":
        return data_missing


@pytest.fixture
def data_repeated(data):
    """
    Generate many datasets.
    Parameters
    ----------
    data : fixture implementing `data`
    Returns
    -------
    Callable[[int], Generator]:
        A callable that takes a `count` argument and
        returns a generator yielding `count` datasets.
    """

    def gen(count):
        for _ in range(count):
            yield data

    return gen


@pytest.fixture
def data_for_sorting():
    """
    Length-3 array with a known sort order.
    This should be three items [B, C, A] with
    A < B < C
    """
    return NullableIntArray(np.array([2, 3, 1]))


@pytest.fixture
def data_missing_for_sorting():
    """
    Length-3 array with a known sort order.
    This should be three items [B, NA, A] with
    A < B and NA missing.
    """
    return NullableIntArray(np.array([2, np.nan, 1]))


@pytest.fixture
def na_cmp():
    """
    Binary operator for comparing NA values.
    Should return a function of two arguments that returns
    True if both arguments are (scalar) NA for your type.
    By default, uses ``operator.is_``
    """
    return operator.is_


@pytest.fixture
def na_value():
    """The scalar missing value for this type. Default 'None'"""
    return np.nan


@pytest.fixture
def data_for_grouping():
    """
    Data for factorization, grouping, and unique tests.
    Expected to be like [B, B, NA, NA, A, A, B, C]
    Where A < B < C and NA is missing
    """
    return NullableIntArray(np.array([2, 2, np.nan, np.nan, 1, 1, 2, 3]))


@pytest.fixture(params=[True, False])
def box_in_series(request):
    """Whether to box the data in a Series"""
    return request.param


@pytest.fixture(
    params=[
        lambda x: 1,
        lambda x: [1] * len(x),
        lambda x: Series([1] * len(x)),
        lambda x: x,
    ],
    ids=["scalar", "list", "series", "object"],
)
def groupby_apply_op(request):
    """
    Functions to test groupby.apply().
    """
    return request.param


@pytest.fixture(params=[True, False])
def as_frame(request):
    """
    Boolean fixture to support Series and Series.to_frame() comparison testing.
    """
    return request.param


@pytest.fixture(params=[True, False])
def as_series(request):
    """
    Boolean fixture to support arr and Series(arr) comparison testing.
    """
    return request.param


@pytest.fixture(params=[True, False])
def use_numpy(request):
    """
    Boolean fixture to support comparison testing of ExtensionDtype array
    and numpy array.
    """
    return request.param


@pytest.fixture(params=["ffill", "bfill"])
def fillna_method(request):
    """
    Parametrized fixture giving method parameters 'ffill' and 'bfill' for
    Series.fillna(method=<method>) testing.
    """
    return request.param


@pytest.fixture(params=[True, False])
def as_array(request):
    """
    Boolean fixture to support ExtensionDtype _from_sequence method testing.
    """
    return request.param


class TestCastingTests(BaseCastingTests):
    pass


class TestConstructorsTests(BaseConstructorsTests):
    pass


class TestDtypeTests(BaseDtypeTests):
    pass


class TestGetitemTests(BaseGetitemTests):
    pass


class TestGroupbyTests(BaseGroupbyTests):
    pass


class TestInterfaceTests(BaseInterfaceTests):
    pass


class TestParsingTests(BaseParsingTests):
    pass


class TestMethodsTests(BaseMethodsTests):
    pass


class TestMissingTests(BaseMissingTests):
    pass


class TestArithmeticOpsTests(BaseArithmeticOpsTests):
    pass


class TestComparisonOpsTests(BaseComparisonOpsTests):
    pass


class TestOpsUtil(BaseOpsUtil):
    pass


class TestUnaryOpsTests(BaseUnaryOpsTests):
    pass


class TestPrintingTests(BasePrintingTests):
    pass


class TestBooleanReduceTests(BaseBooleanReduceTests):
    pass


class TestNoReduceTests(BaseNoReduceTests):
    pass


class TestNumericReduceTests(BaseNumericReduceTests):
    pass


class TestReshapingTests(BaseReshapingTests):
    pass


class TestSetitemTests(BaseSetitemTests):
    pass

内容的提问来源于stack exchange,提问作者Dahn

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.10.06 22:54:02