如何在AzureML流水线的argparse中解析元组参数?
AzureML流水线中解析元组类型参数的解决方案
由于AzureML组件仅支持string、integer、number或bool类型的输入,无法直接传递元组,有两种可行的解决方式:
方案一:通过字符串解析生成元组
1. 修改AzureML组件定义
将输入类型设为string,传递类似"(1,1,1)"或"(1,1,1,12)"格式的字符串:
from azure.ai.ml import command from azure.ai.ml import Input, Output demo_model_training_component = command( name='my sarima pipeline', display_name='my description', description='A long description.', inputs={ "order": Input(type='string'), "seasonal_order": Input(type='string'), }, outputs=dict( df = Output(type="uri_folder", mode="rw_mount") ), code = feature_creation_src_dir, command = """python sarima_model.py \ --order ${{inputs.order}} --seasonal_order ${{inputs.seasonal_order}} \ --df ${{outputs.df}} """, environment = f"{pipeline_job_env.name}{pipeline_job_env.version}", )
2. 修改Python代码中的argparse逻辑
自定义类型解析函数,将传入的字符串转换为元组:
import argparse import statsmodels.api as sm def model_train_sales(X_train, order: tuple, seasonal_order: tuple): model = sm.tsa.SARIMAX(X_train['sales'], order=order, seasonal_order=seasonal_order) results = model.fit() return model, results def parse_tuple(s): # 去除字符串两端的括号,按逗号分割后转为整数元组 s_clean = s.strip('()') return tuple(map(int, s_clean.split(','))) def main(): parser = argparse.ArgumentParser() parser.add_argument("--order", type=parse_tuple) parser.add_argument("--seasonal_order", type=parse_tuple) # 补充读取X_train的逻辑及其他参数 args = parser.parse_args() model, results = model_train_sales(X_train, order=args.order, seasonal_order=args.seasonal_order) if __name__ == "__main__": main()
方案二:拆分参数后组合成元组
将元组的每个元素作为独立参数传入,在代码中重新组合成元组,这种方式类型检查更严格,不易出错。
1. 修改AzureML组件定义
拆分元组元素为独立输入参数:
from azure.ai.ml import command from azure.ai.ml import Input, Output demo_model_training_component = command( name='my sarima pipeline', display_name='my description', description='A long description.', inputs={ "p": Input(type='integer'), "d": Input(type='integer'), "q": Input(type='integer'), "P": Input(type='integer'), "D": Input(type='integer'), "Q": Input(type='integer'), "S": Input(type='integer'), }, outputs=dict( df = Output(type="uri_folder", mode="rw_mount") ), code = feature_creation_src_dir, command = """python sarima_model.py \ --p ${{inputs.p}} --d ${{inputs.d}} --q ${{inputs.q}} \ --P ${{inputs.P}} --D ${{inputs.D}} --Q ${{inputs.Q}} --S ${{inputs.S}} \ --df ${{outputs.df}} """, environment = f"{pipeline_job_env.name}{pipeline_job_env.version}", )
2. 修改Python代码中的argparse逻辑
接收独立参数后组合成元组:
import argparse import statsmodels.api as sm def model_train_sales(X_train, order: tuple, seasonal_order: tuple): model = sm.tsa.SARIMAX(X_train['sales'], order=order, seasonal_order=seasonal_order) results = model.fit() return model, results def main(): parser = argparse.ArgumentParser() parser.add_argument("--p", type=int) parser.add_argument("--d", type=int) parser.add_argument("--q", type=int) parser.add_argument("--P", type=int) parser.add_argument("--D", type=int) parser.add_argument("--Q", type=int) parser.add_argument("--S", type=int) # 补充读取X_train的逻辑及其他参数 args = parser.parse_args() order = (args.p, args.d, args.q) seasonal_order = (args.P, args.D, args.Q, args.S) model, results = model_train_sales(X_train, order=order, seasonal_order=seasonal_order) if __name__ == "__main__": main()
内容的提问来源于stack exchange,提问作者Egorsky
相关产品推荐
相关产品推荐

