Python函数如何复用带类型标注与默认值的公共参数组?
多函数公共参数复用的实现方案
我们可以通过以下几种方案实现公共参数复用,同时保留完整的类型标注、默认值和IDE补全能力,尤其适合机器学习场景下多算法公共参数(学习率、训练轮次、优化器等)的统一管理:
方案1:dataclass封装公共参数(最推荐用于机器学习场景)
将公共参数封装为数据类,默认值、类型标注统一管理,还可扩展参数校验能力:
from dataclasses import dataclass # 公共参数统一定义 @dataclass class CommonTrainParams: b: int = 1 c: str = "hello" # 机器学习场景可直接扩展公共参数: # lr: float = 1e-3 # epoch: int = 100 # optimizer: str = "adam" def f1(a: float = 0.01, common: CommonTrainParams = CommonTrainParams()): b = common.b c = common.c print("do something") def f2(common: CommonTrainParams = CommonTrainParams()): b = common.b c = common.c print("do something different") # 调用示例 f1(a=0.05, common=CommonTrainParams(b=2, c="world")) f2(common=CommonTrainParams(b=3))
- 优势:类型标注完整,IDE全量补全;默认值一处修改全量生效;替换为
pydantic.BaseModel可自动实现参数合法性校验,非常适合调参场景。
方案2:装饰器注入公共参数(保留原生函数调用方式)
如果不需要改动现有函数的调用逻辑,可以通过装饰器自动注入公共参数到函数签名:
import inspect from functools import wraps from typing import Callable # 公共参数模板定义:(参数名, 类型标注, 默认值) COMMON_PARAMS = [ ("b", int, 1), ("c", str, "hello"), ] def inject_common_params(func: Callable) -> Callable: # 合并原函数签名与公共参数 sig = inspect.signature(func) params = list(sig.parameters.values()) for name, anno, default in COMMON_PARAMS: if name not in sig.parameters: params.append(inspect.Parameter( name, inspect.Parameter.KEYWORD_ONLY, annotation=anno, default=default )) func.__signature__ = sig.replace(parameters=params) @wraps(func) def wrapper(*args, **kwargs): return func(*args, **kwargs) return wrapper # 用装饰器修饰函数即可 @inject_common_params def f1(a: float = 0.01, b: int, c: str): print(f"a={a}, b={b}, c={c}") @inject_common_params def f2(b: int, c: str): print(f"b={b}, c={c}") # 调用和原生函数完全一致 f1(a=0.05, b=2, c="world") f2(b=3)
- 优势:无侵入修改原有代码逻辑,调用方式和原生函数完全一致,参数补全、默认值能力完全保留。
方案3:TypedDict+kwargs解构(最低改造成本)
如果习惯用**kwargs传参,可结合TypedDict和Unpack保证类型安全:
from typing import TypedDict, Unpack # 公共参数类型定义 class BCParams(TypedDict, total=False): b: int c: str # 公共参数默认值统一定义 BC_DEFAULTS: BCParams = {"b": 1, "c": "hello"} def f1(a: float = 0.01, **kwargs: Unpack[BCParams]): # 合并默认值与传入参数 params = BC_DEFAULTS | kwargs b, c = params["b"], params["c"] print("do something") def f2(**kwargs: Unpack[BCParams]): params = BC_DEFAULTS | kwargs b, c = params["b"], params["c"] print("do something different") # 调用示例 f1(a=0.05, b=2, c="world") f2(b=3)
- 优势:改造成本极低,不需要调整现有函数的核心逻辑,仅新增类型定义即可实现类型校验。
内容的提问来源于stack exchange,提问作者rkamoi
相关产品推荐
相关产品推荐

