You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

基于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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.13 05:38:10