Pandas ExtensionArray极简实现示例:支持NA的整数数组并通过官方测试
在我看来,Pandas ExtensionArray 是非常需要简单入门示例的场景之一,但我一直没有找到足够简洁的参考案例。
创建 ExtensionArray 的要求
如需创建 ExtensionArray,需要完成以下两项操作:
- 实现
ExtensionDtype并完成注册 - 继承
ExtensionArray并实现所有要求的必要方法
Pandas官方文档的扩展类型章节有相关基础概述。
现有实现参考
目前公开的ExtensionArray实现案例有很多:
- Pandas内置的内部扩展数组
- Geopandas 提供的
GeometryArray - Pandas官方文档中提供的扩展数据类型项目列表
- 例如 CyberPandas 的
IPArray
- 例如 CyberPandas 的
- 其他开源实现,例如 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
相关产品推荐
相关产品推荐

