跳到主要内容

反馈是自然和社会中一切自我调节系统的核心机制。

W. Ross Ashby控制论先驱

🏗️ 数据类

数据类(dataclass)由 dataclasses 模块提供,用于自动生成以数据存储为主的类的样板代码:__init____repr____eq__ 等。当我们需要一个"装数据的容器"时,数据类能比手写类省去大量重复样板,并自带合理的默认行为。

📌 本节要点

  • @dataclass 自动生成 __init____repr____eq__,省去样板代码
  • 字段必须有类型注解,带默认值的字段必须放在无默认值字段后面
  • 可变默认值必须用 field(default_factory=callable),避免共享陷阱
  • field 参数:initreprcomparehashmetadatakw_only 灵活控制字段行为
  • frozen=True 实例不可变且可哈希,slots=True 节省内存,kw_only=True 强制关键字参数
  • __post_init__ 初始化钩子,用于派生字段、参数校验
  • InitVar 仅作为 __init__ 参数传递,不存储为实例属性
  • NamedTuple 对比:NamedTuple 不可变、可解包、更省内存;dataclass 更灵活、可变、可继承 :::
数据类快速体验

基本用法:@dataclass

Python
from dataclasses import dataclass


@dataclass
class FlightState3D:
x: float
y: float
z: float


state = FlightState3D(10.0, 20.0, 100.0)
print(state) # 输出:FlightState3D(x=10.0, y=20.0, z=100.0) ← 自动生成 __repr__
print(state == FlightState3D(10.0, 20.0, 100.0)) # 输出:True ← 自动生成 __eq__
print(state == FlightState3D(10.0, 20.0, 50.0)) # 输出:False

@dataclass 默认会自动生成以下方法(如果未显式定义):

方法默认行为
__init__按声明顺序生成构造器
__repr__形如 FlightState3D(x=10.0, y=20.0, z=100.0)
__eq__按所有字段比较(生成 __eq__,并把 __hash__ 设为 None)
__str__回退到 __repr__
类型注解是必须的

数据类的字段必须有类型注解(如 x: float)。没有注解的赋值会被当作类属性,不进入 __init__

Python
@dataclass
class Bad:
x = 1 # 这是类属性,不是字段!__init__ 不会有 x 参数

字段默认值

带默认值的字段必须放在没有默认值的字段后面(与函数参数规则一致):

Python
from dataclasses import dataclass


@dataclass
class ExperimentConfig:
name: str
model_type: str # 必填参数
learning_rate: float = 0.001 # 默认值
batch_size: int = 32 # 默认值
epochs: int = 100 # 默认值


cfg1 = ExperimentConfig("exp_001", "transformer")
cfg2 = ExperimentConfig("exp_002", "lstm", learning_rate=0.01)
cfg3 = ExperimentConfig("exp_003", "cnn", learning_rate=0.005, batch_size=64, epochs=200)
print(cfg1) # 输出:ExperimentConfig(name='exp_001', model_type='transformer', learning_rate=0.001, batch_size=32, epochs=100)
print(cfg3) # 输出:ExperimentConfig(name='exp_003', model_type='cnn', learning_rate=0.005, batch_size=64, epochs=200)
可变默认值陷阱

不能用可变对象(list、dict、set)作为默认值,否则所有实例会共享同一个对象:

Python
@dataclass
class Bad:
items: list = [] # ValueError: mutable default ... is not allowed

这是 Python 为了避免经典"可变默认参数"陷阱而强制的限制。

field(default_factory=...):可变默认值

field(default_factory=callable) 让每次创建实例时调用 callable() 生成新对象:

Python
from dataclasses import dataclass, field


@dataclass
class ModelConfig:
name: str
hidden_size: int = 128
layer_norm_eps: list[float] = field(default_factory=lambda: [1e-5, 1e-6]) # 每次新建列表
activation_layers: set[str] = field(default_factory=set)

config1 = ModelConfig("encoder")
config2 = ModelConfig("decoder")
config1.activation_layers.add("relu")
config2.activation_layers.add("gelu")
print(config1.activation_layers) # 输出:{'relu'} ← 各自独立
print(config2.activation_layers) # 输出:{'gelu'}

field 的其他参数

Python
from dataclasses import dataclass, field


@dataclass
class SensorConfig:
name: str
sampling_rate: float # Hz
# default_factory:用于可变默认值
calibration_offsets: list[float] = field(default_factory=list)
# repr=False:在 __repr__ 中隐藏(如密钥、校准密钥)
calibration_key: str = field(default="", repr=False)
# compare=False:不参与 __eq__ 比验
firmware_version: str = field(default="1.0.0", compare=False)
# init=False:不进入 __init__,由 __post_init__ 或其他逻辑赋值
display_name: str = field(init=False)

def __post_init__(self) -> None:
# init=False 的字段在这里赋值
self.display_name = f"{self.name} ({self.sampling_rate}Hz)"


s = SensorConfig("IMU", 1000.0, calibration_offsets=[0.01, -0.02], calibration_key="secret")
print(s)
# 输出:SensorConfig(name='IMU', sampling_rate=1000.0, calibration_offsets=[0.01, -0.02], firmware_version='1.0.0', display_name='IMU (1000.0Hz)')
# calibration_key 因 repr=False 被隐藏

s1 = SensorConfig("GPS", 1.0, firmware_version="2.1.0")
s2 = SensorConfig("GPS", 1.0, firmware_version="2.2.0")
print(s1 == s2) # 输出:True(firmware_version 因 compare=False 被忽略)
field 常用参数
  • default:不可变默认值
  • default_factory:可变默认值工厂
  • init(默认 True):是否进入 __init__
  • repr(默认 True):是否出现在 __repr__
  • compare(默认 True):是否参与 __eq__/__lt__
  • hash(默认 None):跟随 compare;设为 True/False 可单独控制
  • metadata:附加元数据字典,供第三方工具使用

@dataclass 的参数

Python
@dataclass(
init=True, # 生成 __init__(默认 True)
repr=True, # 生成 __repr__(默认 True)
eq=True, # 生成 __eq__(默认 True)
order=False, # 生成 __lt__/__le__/__gt__/__ge__(默认 False)
unsafe_hash=False, # 强制生成 __hash__(默认 False)
frozen=False, # 实例不可变(默认 False)
match_args=True, # 生成 __match_args__(3.10+,默认 True)
kw_only=False, # 字段仅关键字参数(3.10+,默认 False)
slots=False, # 生成 __slots__(3.10+,默认 False)
weakref_slot=False, # 添加 __weakref__ 槽位(3.11+,默认 False)
)
class Cfg: ...

order=True:自动生成比较方法

Python
from dataclasses import dataclass


@dataclass(order=True)
class Priority:
level: int
name: str = ""


tasks = [
Priority(3, "低"),
Priority(1, "高"),
Priority(2, "中"),
]
for t in sorted(tasks):
print(t)
# 输出:
# Priority(level=1, name='高')
# Priority(level=2, name='中')
# Priority(level=3, name='低')
order 比较顺序

order=True 按字段声明顺序比较(元组式比较)。如果想让某字段不参与排序,用 field(compare=False)

frozen=True:不可变数据类

Python
from dataclasses import dataclass


@dataclass(frozen=True)
class SimulationParams:
dt: float
duration: float
gravity: float = 9.81


params = SimulationParams(0.01, 10.0)
print(params) # 输出:SimulationParams(dt=0.01, duration=10.0, gravity=9.81)

# params.dt = 0.005 # FrozenInstanceError: cannot assign to field 'dt'

# frozen 数据类自动可哈希(因为 __hash__ 不被设为 None)
print(hash(params)) # 输出:某个整数
print(params in {SimulationParams(0.01, 10.0), SimulationParams(0.02, 5.0)}) # 输出:True
何时用 frozen?
  • 当数据类表示值对象(如仿真参数、坐标、时间戳)时用 frozen=True
  • 需要把实例放入集合、用作字典键时用 frozen=True
  • 多线程/并发场景下,不可变对象更安全
frozen 字段仍可被内部可变对象修改

frozen 只阻止重新赋值字段,不能阻止字段引用的可变对象被修改:

Python
@dataclass(frozen=True)
class Attitude:
roll: float = 0.0
pitch: float = 0.0
yaw: float = 0.0
history: list[float] = field(default_factory=list)

a = Attitude(0.1, 0.2, 0.3)
# a.roll = 0.5 # FrozenInstanceError
a.history.append(1.0) # OK!history 列表本身是可变的

要真正不可变,字段也得是不可变类型(tuple、frozenset 等)。

slots=True(Python 3.10+)

Python
from dataclasses import dataclass


@dataclass(slots=True)
class FlightState3D:
x: float
y: float
z: float


state = FlightState3D(10.0, 20.0, 100.0)
# state.w = 0.0 # AttributeError: 'FlightState3D' object has no attribute 'w'
# 没有 __dict__,节省内存
# print(state.__dict__) # AttributeError

slots=True 自动生成 __slots__,相比手动声明更省事,且与数据类机制完美兼容。

slots 的注意事项
  • slots=True 时不能给实例动态添加属性
  • 子类如果也想要 slots,需要继续声明 @dataclass(slots=True)
  • 想同时支持弱引用,加 weakref_slot=True(Python 3.11+)

kw_only=True(Python 3.10+)

让所有字段只能通过关键字参数传入,避免"位置参数顺序混乱"问题:

Python
from dataclasses import dataclass, field


@dataclass
class ExperimentConfig:
name: str
model_type: str
learning_rate: float = 0.001
batch_size: int = 32

# 传统方式:位置参数容易写错
# ExperimentConfig("exp_001", 32, 0.001, "transformer") ← 把 batch_size 当成 model_type!

@dataclass(kw_only=True)
class ExperimentConfigKW:
name: str
model_type: str
learning_rate: float = 0.001
batch_size: int = 32


# ExperimentConfigKW("exp_001", "transformer") # TypeError: 必须用关键字参数
cfg = ExperimentConfigKW(name="exp_001", model_type="transformer", learning_rate=0.005, batch_size=64)
print(cfg) # 输出:ExperimentConfigKW(name='exp_001', model_type='transformer', learning_rate=0.005, batch_size=64)

单字段 kw_only

也可以只让部分字段成为关键字参数(用 field(kw_only=True)):

Python
from dataclasses import dataclass, field


@dataclass
class FlightController:
name: str
rate: int = 100
timeout: float = field(default=5.0, kw_only=True)
max_retries: int = field(default=3, kw_only=True)


fc = FlightController("imu", 200, timeout=2.0, max_retries=5)
# fc = FlightController("imu", 200, 2.0, 5) # TypeError(后两个必须用关键字)
print(fc)

__post_init__:初始化后的钩子

__init__ 自动生成后,__post_init__ 会在最后被调用,常用于派生字段、参数校验、初始化非字段属性:

Python
from dataclasses import dataclass, field
from datetime import date


@dataclass
class FlightParams:
name: str
takeoff_time: date
# init=False 的派生字段
flight_duration_hours: float = field(init=False)
# 普通字段
created_at: date = field(default_factory=date.today)

def __post_init__(self) -> None:
# 校验
if self.takeoff_time > date.today():
raise ValueError(f"起飞时间不能在未来:{self.takeoff_time}")
# 计算派生字段(示例:假设飞行 2.5 小时)
self.flight_duration_hours = 2.5


fp = FlightParams("maiden_flight", date(2025, 6, 1))
print(fp) # 输出:FlightParams(name='maiden_flight', takeoff_time=datetime.date(2025, 6, 1), flight_duration_hours=2.5, created_at=datetime.date(2025, 7, 14))
print(fp.flight_duration_hours) # 输出:2.5

# FlightParams("future", date(2030, 1, 1)) # ValueError

InitVar:仅用于 __post_init__ 的参数

InitVar 声明的字段会进入 __init__,但不会成为实例属性,只传给 __post_init__

Python
from dataclasses import dataclass, field, InitVar


@dataclass
class HardwareConfig:
name: str
port: int
# 仅初始化参数,不存储为属性
calibration_data: InitVar[list[float]] = field(default_factory=list)

# __post_init__ 接收所有 InitVar
def __post_init__(self, calibration_data: list[float]) -> None:
# 用校准数据计算偏移量,但不保留原始数据
self._offset = sum(calibration_data) / len(calibration_data) if calibration_data else 0.0


hw = HardwareConfig("IMU", 5432, calibration_data=[0.01, -0.02, 0.005])
print(hw) # 输出:HardwareConfig(name='IMU', port=5432)
# calibration_data 不在 repr 中,也不作为属性存在
# print(hw.calibration_data) # AttributeError

NamedTuple 对比

typing.NamedTuple 也能"装数据",且本质是 tuple 子类——不可变、可索引、可解包:

Python
from typing import NamedTuple
from dataclasses import dataclass


class FlightStateNT(NamedTuple):
x: float
y: float
z: float


@dataclass
class FlightStateDC:
x: float
y: float
z: float


# === 共同点 ===
s_nt = FlightStateNT(10.0, 20.0, 100.0)
s_dc = FlightStateDC(10.0, 20.0, 100.0)
print(s_nt) # 输出:FlightStateNT(x=10.0, y=20.0, z=100.0)
print(s_dc) # 输出:FlightStateDC(x=10.0, y=20.0, z=100.0)
print(s_nt == FlightStateNT(10.0, 20.0, 100.0)) # 输出:True
print(s_dc == FlightStateDC(10.0, 20.0, 100.0)) # 输出:True

# === 区别 ===
# 1. NamedTuple 是元组,可索引、可解包
print(s_nt[0]) # 输出:10.0
x, y, z = s_nt # 解包
print(x, y, z) # 输出:10.0 20.0 100.0
# print(s_dc[0]) # AttributeError(dataclass 不支持索引)

# 2. NamedTuple 默认不可变
# s_nt.x = 10.0 # AttributeError
s_dc.x = 50.0 # OK(普通 dataclass 可变)

# 3. NamedTuple 更省内存
import sys
print(sys.getsizeof(s_nt)) # 输出:64 左右
print(sys.getsizeof(s_dc) + sys.getsizeof(s_dc.__dict__)) # 输出:更大
选哪个?
  • NamedTuple:数据是"一组值"、需要不可变、需要解包/索引、内存敏感时(如飞行状态坐标、传感器读数)
  • dataclass:需要可变、需要默认工厂、需要 __post_init__、需要继承、需要更灵活的默认值时
  • frozen dataclass:需要不可变但又不想要 tuple 行为(如仿真参数、配置值对象)

继承与默认值合并

数据类支持继承,子类会合并父类的字段。注意默认值规则:有默认值的字段不能出现在没有默认值的字段之前

Python
from dataclasses import dataclass


@dataclass
class BaseSensor:
name: str


@dataclass
class InheritedSensor(BaseSensor):
sampling_rate: float = 1000.0 # 子类加默认值 OK


sensor = InheritedSensor("accelerometer", 500.0)
print(sensor) # 输出:InheritedSensor(name='accelerometer', sampling_rate=500.0)

# ❌ 反例:父类字段无默认值,子类字段有默认值后,再加无默认值字段会报错
# @dataclass
# class BadSensor(BaseSensor):
# sampling_rate: float = 1000.0
# range: float # TypeError: non-default argument 'range' follows default argument
字段顺序

子类的字段会追加到父类字段之后,整体顺序是"父类字段 + 子类字段"。如果有默认值混排问题,可以用 kw_only=True 解决。

实战:飞行实验配置管理

综合运用 @dataclassfield__post_init__frozenslots,实现一个可序列化、可校验的飞行实验配置:

Python
from dataclasses import dataclass, field, asdict, replace
from typing import Any
import json
import os
import numpy as np


@dataclass
class HardwareConfig:
name: str = "IMU"
port: int = 5432
# 密码不在 repr 中
auth_key: str = field(default="", repr=False)
# 连接池大小
max_connections: int = field(default=10, kw_only=True)

def __post_init__(self) -> None:
if not (1 <= self.port <= 65535):
raise ValueError(f"非法端口:{self.port}")
if self.max_connections < 1:
raise ValueError(f"最大连接数必须为正:{self.max_connections}")

def connection_info(self) -> str:
"""生成连接信息(隐藏密钥)。"""
return f"{self.name}@port:{self.port} (max:{self.max_connections})"


@dataclass
class SimulationConfig:
dt: float = 0.01
duration: float = 10.0
gravity: float = 9.81

def __post_init__(self) -> None:
if self.dt <= 0:
raise ValueError(f"时间步长必须为正:{self.dt}")
if self.duration <= 0:
raise ValueError(f"仿真时长必须为正:{self.duration}")

@property
def steps(self) -> int:
"""仿真总步数。"""
return int(self.duration / self.dt)


@dataclass
class ExperimentConfig:
"""飞行实验配置:聚合硬件、仿真等子配置。"""
name: str = "exp_001"
model_type: str = "transformer"
debug: bool = False
hardware: HardwareConfig = field(default_factory=HardwareConfig)
simulation: SimulationConfig = field(default_factory=SimulationConfig)
# 允许的额外特性开关
features: dict[str, bool] = field(default_factory=lambda: {"real_time": True, "logging": False})

def __post_init__(self) -> None:
if not self.name:
raise ValueError("实验名称不能为空")

def to_json(self, indent: int = 2) -> str:
"""序列化为 JSON 字符串。"""
return json.dumps(asdict(self), ensure_ascii=False, indent=indent)

@classmethod
def from_env(cls) -> "ExperimentConfig":
"""从环境变量加载配置(演示工厂方法)。"""
hw = HardwareConfig(
name=os.getenv("HW_NAME", "IMU"),
port=int(os.getenv("HW_PORT", "5432")),
auth_key=os.getenv("HW_AUTH_KEY", ""),
)
sim = SimulationConfig(
dt=float(os.getenv("SIM_DT", "0.01")),
duration=float(os.getenv("SIM_DURATION", "10.0")),
)
return cls(
name=os.getenv("EXP_NAME", "exp_001"),
model_type=os.getenv("MODEL_TYPE", "transformer"),
debug=os.getenv("DEBUG", "").lower() in ("1", "true", "yes"),
hardware=hw,
simulation=sim,
)

def with_overrides(self, **changes: Any) -> "ExperimentConfig":
"""返回一份修改后的副本(不可变风格)。"""
return replace(self, **changes)


# ===== 使用示例 =====

# 1. 默认配置
config = ExperimentConfig()
print("--- 默认配置 ---")
print(config)
print(config.hardware.connection_info())
# 输出:
# ExperimentConfig(name='exp_001', model_type='transformer', debug=False, hardware=HardwareConfig(name='IMU', ...), ...)
# IMU@port:5432 (max:10)

# 2. 自定义配置
custom = ExperimentConfig(
name="production_test",
model_type="lstm",
debug=False,
hardware=HardwareConfig(name="GPS", port=5432, auth_key="secret", max_connections=20),
simulation=SimulationConfig(dt=0.005, duration=30.0),
)
print("\n--- 自定义配置 ---")
print(custom.hardware.connection_info())
print(f"仿真步数:{custom.simulation.steps}")

# 3. 序列化
print("\n--- JSON 序列化 ---")
print(custom.to_json())
# 输出:完整的 JSON 字符串(含嵌套结构)

# 4. 校验
print("\n--- 校验 ---")
try:
bad = HardwareConfig(port=99999)
except ValueError as e:
print(f"拦截:{e}") # 输出:拦截:非法端口:99999

# 5. 不可变风格的修改:replace
print("\n--- 不可变修改(replace)---")
dev_config = custom.with_overrides(debug=True, name="dev_test")
print(f"原配置:debug={custom.debug}, name={custom.name}")
print(f"新配置:debug={dev_config.debug}, name={dev_config.name}")
# replace 返回新对象,原对象不变

# 6. 从环境变量加载
print("\n--- 环境变量加载 ---")
os.environ["HW_NAME"] = "LIDAR"
os.environ["HW_PORT"] = "6543"
os.environ["DEBUG"] = "true"
env_config = ExperimentConfig.from_env()
print(env_config.hardware.name, env_config.hardware.port) # 输出:LIDAR 6543
print(env_config.debug) # 输出:True
常用工具函数
  • asdict(obj):把数据类实例递归转为字典(适合 JSON 序列化)
  • astuple(obj):转为元组
  • replace(obj, **changes):返回修改部分字段后的新实例(不可变修改)
  • fields(obj) / fields(Cls):返回字段信息列表
  • is_dataclass(obj):判断是否为数据类实例或类型

frozen + slots 的现代配置

Python
from dataclasses import dataclass, field
import numpy as np


@dataclass(frozen=True, slots=True)
class FlightControllerConfig:
"""不可变 + 节省内存的现代控制器配置。"""
name: str
loop_rate: int = 100
timeout: float = 5.0
gains: dict[str, float] = field(default_factory=dict)

def __post_init__(self) -> None:
# frozen 时不能直接 self.loop_rate = ...
# 需要修改要用 object.__setattr__
if not self.gains:
object.__setattr__(self, "gains", {"kp": 1.0, "ki": 0.1, "kd": 0.01})

fc = FlightControllerConfig("attitude_controller")
print(fc.gains) # 输出:{'kp': 1.0, 'ki': 0.1, 'kd': 0.01}
# fc.loop_rate = 200 # FrozenInstanceError
frozen 中修改字段

frozen=True 时普通赋值会抛 FrozenInstanceError。在 __post_init__ 中需要修改字段时,用 object.__setattr__(self, name, value) 绕过冻结检查(这是合法用法,标准库也这么做)。

🎯 动手练习

  1. 不可变配置:使用 frozen=Trueslots=True 创建 SimulationParams 数据类,包含 dt、duration、gravity 字段,在 __post_init__ 中校验 dt > 0
  2. 传感器系统:实现 SensorSuite 数据类,包含传感器列表(default_factory=list)、采样率(init=False 派生字段),在 __post_init__ 中计算平均采样率
  3. 实验响应:创建 ExperimentResult 数据类,使用 field(repr=False) 隐藏敏感数据,实现 to_dict()to_json() 方法
  4. NamedTuple vs dataclass:分别用两种方式实现 FlightState 类,对比内存占用、可变性、解包能力

📚 延伸阅读

  • Pydantic:第三方库,提供运行时数据校验和类型强制转换
  • attrsdataclasses 的前身,功能更丰富(转换器、验证器、槽位等)
  • TypedDicttyping.TypedDict 用于字典结构的类型注解
  • 模式匹配:Python 3.10+ 的 match-case 与数据类的解构匹配
📋速查表
基本数据类`@dataclass`
@dataclass class FlightState3D: x: float; y: float; z: float
默认值`field: type = value`
learning_rate: float = 0.001
可变默认值`field(default_factory=...)`
gains: dict = field(default_factory=dict)
隐藏 repr`field(repr=False)`
auth_key: str = field(default="", repr=False)
不参与比较`field(compare=False)`
firmware_version: str = field(default="1.0", compare=False)
不进入 init`field(init=False)`
display_name: str = field(init=False)
不可变`@dataclass(frozen=True)`
实例不可修改,自动可哈希
节省内存`@dataclass(slots=True)`
生成 __slots__,3.10+
关键字参数`@dataclass(kw_only=True)`
字段只能用关键字传入,3.10+
生成比较`@dataclass(order=True)`
生成 __lt__、__le__ 等
初始化钩子`def __post_init__(self):`
派生字段、校验参数
仅初始化参数`InitVar[type]`
传入 __init__ 但不存储
转字典`asdict(obj)`
递归转为嵌套字典
修改副本`replace(obj, **changes)`
返回修改后的新实例

✅ 本节总结

本节我们学习了数据类,核心要点包括:

  • @dataclass 自动生成样板代码__init____repr____eq__ 一键生成,省去重复代码
  • 字段必须有类型注解:这是数据类识别字段的依据,无注解的赋值会被当作类属性
  • 可变默认值用 default_factory:避免所有实例共享同一个可变对象的经典陷阱
  • field 参数灵活控制initreprcomparehash 精细控制每个字段的行为
  • frozen=True 实现不可变:适合值对象、配置类,自动可哈希可用作字典键
  • slots=True 节省内存:3.10+ 特性,与数据类机制完美兼容
  • __post_init__ 是初始化钩子:用于派生字段、参数校验、初始化 init=False 的字段
  • InitVar 传递参数不存储:适合校准数据等只需初始化时使用但不保留的数据
  • NamedTuple vs dataclass:前者是不可变元组、可解包、更省内存;后者更灵活、可变、可继承

至此,面向对象章节的内容告一段落。综合运用类、继承、多态、封装、魔术方法和数据类,已能写出结构清晰、行为地道、可维护性强的 Python 程序了。