卖兜搞IT

2024彻底搞懂Python的Dataclass(2)

本文较长,会分为几期,目标是给大家讲解Python的dataclass该怎么用,以及它是如何工作的,而不只是和大家一起创建一个dataclass

本文基于Python的版本是3.11,不同的版本的python,dataclass可能会有不同的差异,这是因为dataclass是一个正在不断演进的python特性,变化较大。

本文为连载,如果没有看过前面的内容,建议先跳转下面的链接进行“补课”

2024彻底搞懂Python的Dataclass(1)

__eq__  比较

dataclass不光帮我们默认实现了__str__和__repr__方法,而且连__eq__ 默认也实现了,而且比较的原则是已实例的属性值为依据,而不是传统class默认的以实例的ID为依据

@dataclass
class Position:
    x: float = 0
    y: float = 0
>>> p1 = Position(x=1, y=-2)
>>> p2 = Position(x=1, y=-2)
>>> p1 == p2
True

传统class要实现这一的比较,需要我们自己实现__eq__方法,如下:

class Position:
    def __init__(self, x: float = 0, y: float = 0,):
        self.x = x
        self.y = y
    def __str__(self):
        return f"{self.__class__.__qualname__}(x={self.x}, y={self.y})"
    def __repr__(self):
        return f"{self.__class__.__qualname__}(x={self.x}, y={self.y})"

默认比较是比较两个实例的ID,所以肯定是不同的

>>> p1 = Position(x=1, y=-2)
>>> p2 = Position(x=1, y=-2)
>>> p1 == p2
False

添加__eq__方法后,才实现dataclass的效果

class Position:
    def __init__(self, x: float = 0, y: float = 0,):
        self.x = x
        self.y = y
    def __str__(self):
        return f"{self.__class__.__qualname__}(x={self.x}, y={self.y})"
    def __repr__(self):
        return f"{self.__class__.__qualname__}(x={self.x}, y={self.y})"
    def __eq__(self, other):
        if self.__class__ == other.__class__:
            return (self.x, self.y) == (other.x, other.y)
        return NotImplemented
>>> p1 = Position(x=1, y=-2)
>>> p2 = Position(x=1, y=-2)
>>> p1 == p2
True

p1 == p2和p1 is p2 不是一回事,请大家接着看。

__hash__ 方法

如果我们手动添加__eq__方法,那么__hash__ 方法我们也需要重写,为啥呢?

what is __hash__

首先我们先解释了__hash__ 方法是干啥的。每一个python object都有一个属性 __hash__

>>> dir(object)
['__class__',
 '__delattr__',
 '__dir__',
 '__doc__',
 '__eq__',
 '__format__',
 '__ge__',
 '__getattribute__',
 '__getstate__',
 '__gt__',
 '__hash__',
 '__init__',
 '__init_subclass__',
 '__le__',
 '__lt__',
 '__ne__',
 '__new__',
 '__reduce__',
 '__reduce_ex__',
 '__repr__',
 '__setattr__',
 '__sizeof__',
 '__str__',
 '__subclasshook__']

但并不代表每一个python object都是可哈希的。比如一个list,就不是可哈希的,因为list是可变数据类型,它的__hash__属性是None

>>> a = [1, 2]
>>> print(a.__hash__)
None
>>> hash(a)
Traceback (most recent call last):
  File "<stdin>", line 1, in <module>
TypeError: unhashable type: 'list'

如果一个python object不可哈希,它也就不能作为python dict的key,或者set的元素。

>>> a = [1, 2]
>>> b = {a: 1}
Traceback (most recent call last):
  File "<stdin>", line 1, in <module>
TypeError: unhashable type: 'list'

Python默认class是可哈希的

class Position:
    def __init__(self, x: float = 0, y: float = 0,):
        self.x = x
        self.y = y
>>> p1 = Position(1,1)
>>> p2 = Position(1,1)
>>> p1 == p2
False
>>>> p1 is p2
False
>>> hash(p1), hash(p2)  # 对各自对象所在的内容地址进行hash
(17592085467597, 17592085467665)
>>> a = {p1:p2}  # 可哈希,所以就可以作为dict的key

为啥class添加了__eq__就不可哈希了呢?

一句话总结,就是出现了矛盾,当我们添加了我们的__eq__方法后,

class Position:
    def __init__(self, x: float = 0, y: float = 0,):
        self.x = x
        self.y = y
    def __eq__(self, other):
        if self.__class__ == other.__class__:
            return (self.x, self.y) == (other.x, other.y)
        return NotImplemented

p1是等于p2了,但是p1和p2的内存地址还是不同的,导致两个对象如果按照内存地址算hash,hash就不同了。

>>> p1 == p2
True
>>> p1 is p2
False

为避免出现这个矛盾,所以干脆就让这个class的__hash___是None,也就是不可哈希了

>>> hash(p1)
Traceback (most recent call last):
  File "<stdin>", line 1, in <module>
TypeError: unhashable type: 'Position'

所以回到前面,当我们重写了__eq__方法后,__hash___方法也需要重写。特别是,我们如果需要我们的class是可哈希的话。简单重写如下:

class Position:
    def __init__(self, x: float = 0, y: float = 0,):
        self.x = x
        self.y = y
    def __str__(self):
        return f"{self.__class__.__qualname__}(x={self.x}, y={self.y})"
    def __repr__(self):
        return f"{self.__class__.__qualname__}(x={self.x}, y={self.y})"

            def __eq__(self, other):
        if self.__class__ == other.__class__:
            return (self.x, self.y) == (other.x, other.y)
        return NotImplemented
    def __hash__(self):
        return hash((self.x, self.y))

这时,它们就是可哈希的,而且我们用x和y作为一个元祖值去计算哈希,所以当两个对象的x和y相同时,它们的哈希值也一样。

>>> p1 = Position(1,1)
>>> p2 = Position(1,1)
>>> p1 == p2
True
>>> p1 is p2
False
>>> hash(p1), hash(p2)
(8389048192121911274, 8389048192121911274)

当然也可以作为dict的key出现,因为p1和p2的哈希一样,所以就是重复的key

>>> p1 = Position(1,1)
>>> p2 = Position(1,1)
>>>
>>> d = {p1: 1, p2: 2}
>>> d
{Position(x=1, y=1): 2}

还有一个潜在问题

一般来说,可哈希的对象,一般是immutable的。但是我们这个position对象,却可以改变它的x和y值

比如我们可以改变p1的x值,如此又导致了各种矛盾问题。

>>> p1.x
1
>>> p1.x = 2
>>> d
{Position(x=2, y=1): 2}
>>> p1 == p2
False
>>> hash(p1), hash(p2)
(6794810172467074373, 8389048192121911274)
>>>

你还可以把p1的值再改回去,等等会出现各种奇怪现象,那为了尽量避免这个现象,我们需要引入property来保护一下我们的x和y

class Position:
    def __init__(self, x: float = 0, y: float = 0,):
        self._x = x
        self._y = y
    @property
    def x(self):
        return self._x
    @property
    def y(self):
        return self._y
    def __str__(self):
        return f"{self.__class__.__qualname__}(x={self.x}, y={self.y})"
    def __repr__(self):
        return f"{self.__class__.__qualname__}(x={self.x}, y={self.y})"

            def __eq__(self, other):
        if self.__class__ == other.__class__:
            return (self.x, self.y) == (other.x, other.y)
        return NotImplemented
    def __hash__(self):
        return hash((self.x, self.y))

>>> p1 = Position(1,1)
>>> p2 = Position(1,1)
>>> p1 == p2
True
>>> d = {p1: 1}
>>> p1.x
1
>>> p1.x = 2
Traceback (most recent call last):
  File "<stdin>", line 1, in <module>
AttributeError: property 'x' of 'Position' object has no setter

OK,上面我们花了这么大的篇幅,来实现一个Position类,而且很容易出错,或者忘记,但是如果用dataclass,就可以非常简单了,只需要下面几行代码就可以实现我们上面一个包含7个def的class

用dataclass四行代码搞定

from dataclasses import dataclass

@dataclass(frozen=True)
class Position:
    x: float = 0
    y: float = 0

测试一下,所有之前需要__init__, property, __eq__, __str__, __repr__, __hash__等 ,全部省略,dataclass都帮我们做了。

>>> p1 = Position(1, 1)
>>> p2 = Position(1, 1)
>>> p1 == p2
True
>>> hash(p1), hash(p2)
(8389048192121911274, 8389048192121911274)
>>>
>>> d = {p1: 1}
>>> p1.x
1
>>> p1.x = 2
Traceback (most recent call last):
  File "<stdin>", line 1, in <module>
  File "<string>", line 4, in __setattr__
dataclasses.FrozenInstanceError: cannot assign to field 'x'
>>>

~~~~~~✍️未完待续✍️~~~~~~