反馈是自然和社会中一切自我调节系统的核心机制。
🏗️ 数据类
数据类(dataclass)由 dataclasses 模块提供,用于自动生成以数据存储为主的类的样板代码:__init__、__repr__、__eq__ 等。当我们需要一个"装数据的容器"时,数据类能比手写类省去大量重复样板,并自带合理的默认行为。
📌 本节要点
@dataclass自动生成__init__、__repr__、__eq__,省去样板代码- 字段必须有类型注解,带默认值的字段必须放在无默认值字段后面
- 可变默认值必须用
field(default_factory=callable),避免共享陷阱 field参数:init、repr、compare、hash、metadata、kw_only灵活控制字段行为frozen=True实例不可变且可哈希,slots=True节省内存,kw_only=True强制关键字参数__post_init__初始化钩子,用于派生字段、参数校验InitVar仅作为__init__参数传递,不存储为实例属性- 与
NamedTuple对比:NamedTuple 不可变、可解包、更省内存;dataclass 更灵活、可变、可继承 :::
基本用法:@dataclass
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__:
@dataclass
class Bad:
x = 1 # 这是类属性,不是字段!__init__ 不会有 x 参数
字段默认值
带默认值的字段必须放在没有默认值的字段后面(与函数参数规则一致):
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)作为默认值,否则所有实例会共享同一个对象:
@dataclass
class Bad:
items: list = [] # ValueError: mutable default ... is not allowed
这是 Python 为了避免经典"可变默认参数"陷阱而强制的限制。
field(default_factory=...):可变默认值
用 field(default_factory=callable) 让每次创建实例时调用 callable() 生成新对象:
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 的其他参数
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 被忽略)
default:不可变默认值default_factory:可变默认值工厂init(默认 True):是否进入__init__repr(默认 True):是否出现在__repr__compare(默认 True):是否参与__eq__/__lt__hash(默认 None):跟随 compare;设为 True/False 可单独控制metadata:附加元数据字典,供第三方工具使用
@dataclass 的参数
@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:自动生成比较方法
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=True 按字段声明顺序比较(元组式比较)。如果想让某字段不参与排序,用 field(compare=False)。
frozen=True:不可变数据类
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=True - 需要把实例放入集合、用作字典键时用
frozen=True - 多线程/并发场景下,不可变对象更安全
frozen 只阻止重新赋值字段,不能阻止字段引用的可变对象被修改:
@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+)
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=True时不能给实例动态添加属性- 子类如果也想要 slots,需要继续声明
@dataclass(slots=True) - 想同时支持弱引用,加
weakref_slot=True(Python 3.11+)
kw_only=True(Python 3.10+)
让所有字段只能通过关键字参数传入,避免"位置参数顺序混乱"问题:
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)):
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__ 会在最后被调用,常用于派生字段、参数校验、初始化非字段属性:
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__:
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 子类——不可变、可索引、可解包:
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 行为(如仿真参数、配置值对象)
继承与默认值合并
数据类支持继承,子类会合并父类的字段。注意默认值规则:有默认值的字段不能出现在没有默认值的字段之前:
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 解决。
实战:飞行实验配置管理
综合运用 @dataclass、field、__post_init__、frozen、slots,实现一个可序列化、可校验的飞行实验配置:
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 的现代配置
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=True 时普通赋值会抛 FrozenInstanceError。在 __post_init__ 中需要修改字段时,用 object.__setattr__(self, name, value) 绕过冻结检查(这是合法用法,标准库也这么做)。
🎯 动手练习
- 不可变配置:使用
frozen=True和slots=True创建SimulationParams数据类,包含 dt、duration、gravity 字段,在__post_init__中校验 dt > 0 - 传感器系统:实现
SensorSuite数据类,包含传感器列表(default_factory=list)、采样率(init=False派生字段),在__post_init__中计算平均采样率 - 实验响应:创建
ExperimentResult数据类,使用field(repr=False)隐藏敏感数据,实现to_dict()和to_json()方法 - NamedTuple vs dataclass:分别用两种方式实现
FlightState类,对比内存占用、可变性、解包能力
📚 延伸阅读
- Pydantic:第三方库,提供运行时数据校验和类型强制转换
- attrs:
dataclasses的前身,功能更丰富(转换器、验证器、槽位等) - TypedDict:
typing.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参数灵活控制:init、repr、compare、hash精细控制每个字段的行为frozen=True实现不可变:适合值对象、配置类,自动可哈希可用作字典键slots=True节省内存:3.10+ 特性,与数据类机制完美兼容__post_init__是初始化钩子:用于派生字段、参数校验、初始化init=False的字段InitVar传递参数不存储:适合校准数据等只需初始化时使用但不保留的数据- NamedTuple vs dataclass:前者是不可变元组、可解包、更省内存;后者更灵活、可变、可继承
至此,面向对象章节的内容告一段落。综合运用类、继承、多态、封装、魔术方法和数据类,已能写出结构清晰、行为地道、可维护性强的 Python 程序了。