Jupyter导入utilities.py的Sagemaker模型加载函数执行shell命令报错如何解决
错误原因
!是Jupyter Notebook专属的行魔法命令,仅在Notebook单元格直接运行时会被IPython内核解析为shell命令执行。当你把带!的命令写在普通.py工具文件中、通过exec执行时,原生Python解释器不识别!语法,因此抛出语法错误。
解决方法
不需要把函数复制到每个Notebook中,直接修改utilities.py里的函数,替换掉exec+!的魔法命令写法即可,下面提供两种可选实现方案:
方案1:用subprocess执行aws cli命令
适合已经在环境中安装配置好aws cli的场景,修改后的函数如下:
def get_aws_sagemaker_model(model_loc): """ TO BE USED IN A JUPYTER NOTEBOOK extracts a sagemaker model that has ran and been completed deletes the copied items and leaves you with the model note that you will need to have the package installed with correct versioning for whatever model you have trained ie. if you are loading an XGBoost model, have XGBoost installed Args: model_loc (str) : s3 location of the model including file name Return: model: unpacked and loaded model """ import re import tarfile import os import pickle as pkl import subprocess # extract the filename from beyond the last backslash packed_model_name = re.search("(.*\/)(.*)$" , model_loc)[2] # 替换原来的exec魔法命令,用subprocess执行aws cli subprocess.run(["aws", "s3", "cp", model_loc, "."], check=True) # use tarfile to extract tar = tarfile.open(packed_model_name) # extract filename from tarfile unpacked_model_name = tar.getnames()[0] tar.extractall() tar.close() model = pkl.load(open(unpacked_model_name, 'rb')) # cleanup copied files and unpacked model os.remove(packed_model_name) os.remove(unpacked_model_name) return model
方案2:用boto3直接下载S3文件(更推荐)
不需要依赖本地aws cli,纯Python实现兼容性更强,需要先在环境中安装boto3(执行pip install boto3即可),修改后的函数如下:
def get_aws_sagemaker_model(model_loc): """ TO BE USED IN A JUPYTER NOTEBOOK extracts a sagemaker model that has ran and been completed deletes the copied items and leaves you with the model note that you will need to have the package installed with correct versioning for whatever model you have trained ie. if you are loading an XGBoost model, have XGBoost installed Args: model_loc (str) : s3 location of the model including file name Return: model: unpacked and loaded model """ import re import tarfile import os import pickle as pkl import boto3 from urllib.parse import urlparse # extract the filename from beyond the last backslash packed_model_name = re.search("(.*\/)(.*)$" , model_loc)[2] # 用boto3直接下载S3文件,不需要aws cli parsed_s3_path = urlparse(model_loc) bucket_name = parsed_s3_path.netloc file_key = parsed_s3_path.path.lstrip('/') s3_client = boto3.client('s3') s3_client.download_file(bucket_name, file_key, packed_model_name) # use tarfile to extract tar = tarfile.open(packed_model_name) # extract filename from tarfile unpacked_model_name = tar.getnames()[0] tar.extractall() tar.close() model = pkl.load(open(unpacked_model_name, 'rb')) # cleanup copied files and unpacked model os.remove(packed_model_name) os.remove(unpacked_model_name) return model
注意事项
- 两种方案都需要确保当前执行环境的IAM角色/访问凭证有对应S3路径的读权限
- pickle反序列化存在安全风险,仅加载你自己信任的模型文件
- 如果加载的是XGBoost、PyTorch等框架的模型,也可以直接用对应框架的S3读取接口,不需要先下载到本地
内容的提问来源于stack exchange,提问作者ESlice
相关产品推荐
相关产品推荐

