基于Python实现类构造函数的参数解析(仿MATLAB inputParser)
问题描述
我是Python新手,想创建一个类的构造函数,要求包含少量必填属性,以及大量带默认值和输入校验规则的可选属性。我试过用argparse模块,但不知道怎么把解析后的参数传入类属性,而且这个模块没法定义输入逻辑校验规则。希望实现类似下面MATLAB脚本的功能:
methods function obj = Platform(ClassID,varargin) inPar = inputParser; expectedClass = {'Ownship', 'Wingman', 'Flight Group', 'Unknown', 'Suspect', 'Neutral', 'Friend', 'Foe'}; validClassID = @(x) any(validatestring(x,expectedClass)); addRequired(inPar,'ClassID',validClassID) defaultDim = struct('length', 0, 'width', 0, 'height', 0, 'oOffset', [0 0 0]); validDim = @(x) ~isempty(intersect(fieldnames(x),fieldnames(defaultDim))); addOptional(inPar,'Dimensions',defaultDim,validDim) defaultPos = [0 0 0]; validPos = @(x) isclass(x,'double') && mean(size(x) == [1 3]); addOptional(inPar,'Position',defaultPos,validPos) defaultOr = [0 0 0]; validOr = @(x) isclass(x,'double') && mean(size(x) == [1 3]); addOptional(inPar,'Orientation',defaultOr,validOr) defaultTraj = struct('Waypoints',[0 0 0],... 'TimeofArrival',0,... 'Velocity',[0 0 0],... 'Orientation',[0 0 0]); validTraj = @(x) ~isempty(fieldnames(x),fieldnames(defaultTraj)); addOptional(inPar,'Trajectory',defaultTraj,validTraj) expectedDL = {'One','Two','Three'}; defaultDL = {}; validDL = @(x) any(validatestring(x,expectedDL)); addOptional(inPar,'DataLinks',defaultDL,validDL) defaultSens = {}; validSens = @(x) isa(x,'Sensor'); addOptional(inPar,'Sensors',defaultSens,validSens) parse(inPar,ClassID,varargin{:}) obj.PlatformID = randi([1 10000]); obj.ClassID = inPar.Results.ClassID; obj.Dimensions = inPar.Results.Dimensions; obj.Position = inPar.Results.Position; obj.Orientation = inPar.Results.Orientation; obj.Trajectory = inPar.Results.Trajectory; obj.Sensors = inPar.Results.Sensors; obj.DataLinks = inPar.Results.DataLinks; end
解决方案
方法一:纯Python手动实现(无需第三方库)
直接在类的__init__方法中处理必填参数校验,用**kwargs接收可选参数,配合自定义校验函数和默认值,完全匹配MATLAB的inputParser逻辑:
import random from typing import List, Dict, Union import numpy as np # 定义Sensor类(对应MATLAB中的Sensor类) class Sensor: pass class Platform: def __init__(self, class_id: str, **kwargs): # 校验必填参数ClassID expected_classes = {'Ownship', 'Wingman', 'Flight Group', 'Unknown', 'Suspect', 'Neutral', 'Friend', 'Foe'} if class_id not in expected_classes: raise ValueError(f"ClassID必须是以下值之一: {expected_classes}") # 定义可选参数的默认值与校验规则 default_dim = {'length': 0, 'width': 0, 'height': 0, 'oOffset': np.array([0, 0, 0])} def validate_dim(dim: Dict) -> bool: return bool(set(dim.keys()) & set(default_dim.keys())) default_pos = np.array([0, 0, 0]) def validate_pos(pos: np.ndarray) -> bool: return isinstance(pos, np.ndarray) and pos.shape in [(1, 3), (3,)] default_or = np.array([0, 0, 0]) def validate_or(or_val: np.ndarray) -> bool: return isinstance(or_val, np.ndarray) and or_val.shape in [(1, 3), (3,)] default_traj = {'Waypoints': np.array([0, 0, 0]), 'TimeofArrival': 0, 'Velocity': np.array([0, 0, 0]), 'Orientation': np.array([0, 0, 0])} def validate_traj(traj: Dict) -> bool: return bool(set(traj.keys()) & set(default_traj.keys())) expected_dl = {'One', 'Two', 'Three'} default_dl = [] def validate_dl(dl: Union[str, List[str]]) -> bool: if isinstance(dl, str): return dl in expected_dl return all(item in expected_dl for item in dl) default_sens = [] def validate_sens(sens: List[Sensor]) -> bool: return all(isinstance(item, Sensor) for item in sens) # 解析参数并赋值给类属性 self.platform_id = random.randint(1, 10000) self.class_id = class_id self.dimensions = self._get_validated_param('Dimensions', kwargs, default_dim, validate_dim) self.position = self._get_validated_param('Position', kwargs, default_pos, validate_pos) self.orientation = self._get_validated_param('Orientation', kwargs, default_or, validate_or) self.trajectory = self._get_validated_param('Trajectory', kwargs, default_traj, validate_traj) self.data_links = self._get_validated_param('DataLinks', kwargs, default_dl, validate_dl) self.sensors = self._get_validated_param('Sensors', kwargs, default_sens, validate_sens) def _get_validated_param(self, param_name: str, kwargs: Dict, default: any, validator: callable) -> any: """辅助函数:统一处理参数获取与校验""" value = kwargs.get(param_name, default) if not validator(value): raise ValueError(f"参数{param_name}不符合校验规则") return value
方法二:使用Pydantic库(更简洁高效)
如果允许使用第三方库,pydantic专门用于数据校验和默认值管理,代码更简洁易维护:
首先安装依赖:
pip install pydantic
实现代码:
import random from typing import List, Dict, Union import numpy as np from pydantic import BaseModel, validator, Field class Sensor: pass class PlatformConfig(BaseModel): # 必填参数 class_id: str = Field(..., description="平台类型") # 可选参数及默认值 dimensions: Dict[str, Union[int, np.ndarray]] = Field(default={'length': 0, 'width': 0, 'height': 0, 'oOffset': np.array([0,0,0])}) position: np.ndarray = Field(default=np.array([0,0,0])) orientation: np.ndarray = Field(default=np.array([0,0,0])) trajectory: Dict[str, Union[int, np.ndarray]] = Field(default={'Waypoints': np.array([0,0,0]), 'TimeofArrival':0, 'Velocity':np.array([0,0,0]), 'Orientation':np.array([0,0,0])}) data_links: List[str] = Field(default=[]) sensors: List[Sensor] = Field(default=[]) # 校验规则定义 @validator('class_id') def check_class_id(cls, v): expected_classes = {'Ownship', 'Wingman', 'Flight Group', 'Unknown', 'Suspect', 'Neutral', 'Friend', 'Foe'} if v not in expected_classes: raise ValueError(f"必须是以下值之一: {expected_classes}") return v @validator('dimensions') def check_dimensions(cls, v): default_keys = {'length', 'width', 'height', 'oOffset'} if not set(v.keys()) & default_keys: raise ValueError("Dimensions必须包含至少一个默认字段") return v @validator('position', 'orientation') def check_3d_array(cls, v): if not isinstance(v, np.ndarray) or v.shape not in [(1,3), (3,)]: raise ValueError("必须是形状为(1,3)或(3,)的numpy数组") return v @validator('trajectory') def check_trajectory(cls, v): default_keys = {'Waypoints', 'TimeofArrival', 'Velocity', 'Orientation'} if not set(v.keys()) & default_keys: raise ValueError("Trajectory必须包含至少一个默认字段") return v @validator('data_links', each_item=True) def check_data_links(cls, v): expected_dl = {'One', 'Two', 'Three'} if v not in expected_dl: raise ValueError(f"必须是以下值之一: {expected_dl}") return v @validator('sensors', each_item=True) def check_sensors(cls, v): if not isinstance(v, Sensor): raise ValueError("必须是Sensor类的实例") return v class Platform: def __init__(self, class_id: str, **kwargs): config = PlatformConfig(class_id=class_id, **kwargs) self.platform_id = random.randint(1, 10000) self.class_id = config.class_id self.dimensions = config.dimensions self.position = config.position self.orientation = config.orientation self.trajectory = config.trajectory self.data_links = config.data_links self.sensors = config.sensors
内容的提问来源于stack exchange,提问作者KerbFusion
相关产品推荐
相关产品推荐

