如何在app.py中导入带argparse命令行参数的train.py并设置参数
解决方案
分为两种场景,优先选择可修改train.py的方案,稳定性更强。
方案1:修改train.py结构(推荐)
把train.py的原有执行逻辑封装为可接收参数的函数,兼顾命令行调用和模块导入调用两种场景:
修改后的train.py示例
import argparse # 其他原有导入语句 model_names = ["resnet18", "resnet50", ...] # 你的原有模型列表定义 # 把参数解析逻辑单独拆为函数,方便外部调用 def get_parser(): parser = argparse.ArgumentParser(description='PyTorch ImageNet Training') parser.add_argument('data', metavar='DIR', help='path to dataset') parser.add_argument('--model-dir', type=str, default='', help='path to desired output directory for saving model ' 'checkpoints (default: current directory)') parser.add_argument('-a', '--arch', metavar='ARCH', default='resnet18', choices=model_names, help='model architecture: ' + ' | '.join(model_names) + ' (default: resnet18)') # 其余所有add_argument参数定义都放在这里 return parser # 把原有训练逻辑全部放到main函数里 def main(args=None): parser = get_parser() # 没有传入参数时走命令行解析逻辑 if args is None: args = parser.parse_args() # 以下是你原来的所有训练代码,直接使用args变量即可 print(f"数据集路径:{args.data},模型保存目录:{args.model_dir},模型架构:{args.arch}") # 剩余训练逻辑... if __name__ == "__main__": main()
app.py调用方式
有两种传参方式可选:
- 方式1:传入命令行格式的参数列表
import train # 列表里的参数顺序和命令行调用的顺序完全一致即可 train.main([ "path_to_datafolder", "--model-dir=sdcsdc", "--batch-size=333", "--arch=resnet50" ])
- 方式2:构造Namespace对象传参,更灵活可控
import train from argparse import Namespace # 先获取默认参数再按需修改,避免漏写参数 parser = train.get_parser() args = parser.parse_args([]) # 拿到所有参数的默认值 # 自定义修改需要的参数 args.data = "path_to_datafolder" args.model_dir = "sdcsdc" args.batch_size = 333 args.arch = "resnet50" # 传入参数执行训练 train.main(args)
方案2:不修改train.py的临时方案
如果无法修改train.py的代码,可以通过修改sys.argv模拟命令行参数,导入时直接执行train.py的逻辑:
import sys # 按命令行格式设置argv,第一个元素固定为脚本名,后面跟参数 sys.argv = [ "train.py", "path_to_datafolder", "--model-dir=sdcsdc", "--batch-size=333" ] # 导入train时会自动执行其中的代码,使用上面设置的参数 import train
注意:该方案仅适用于train.py没有
if __name__ == "__main__"判断的场景,且train的代码只会在第一次导入时执行,复用性较差,仅作为临时方案使用。
内容的提问来源于stack exchange,提问作者Ilgar Rasulov
相关产品推荐
相关产品推荐

