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

如何在Amazon SageMaker中将预测图像保存至S3并返回图像?

问题

我在inference.py中编写了如下output_fn函数:

def output_fn(prediction, content_type):
    assert content_type == 'application/json'
    mask = prediction['mask']
    mask = np.where(mask==2, 255, mask)
    mask = np.where(mask==1, 128, mask)
    trimap = mask.astype(np.uint8)
    trimap_image = Image.fromarray(trimap).resize((prediction['w'], prediction['h'])).convert('L')
    trimap_image.save(prediction['save_path'])
    return json.dumps(['Image saved!'])

当前我将prediction['save_path']设为'test.png',但无法找到该图像。请问如何修改prediction['save_path']以将图像保存至S3桶?另外,是否有办法直接返回PNG图像或Numpy数组?我的模型部署代码如下:

predictor = model.deploy(
    initial_instance_count=1,
    instance_type=instance_type,
    serializer=JSONSerializer(),
    deserializer=JSONDeserializer(),   
)

一、将图像保存到S3桶

Image.save()无法直接写入S3路径,需要借助AWS的boto3库完成上传,步骤如下:

  1. 确保你的SageMaker执行角色拥有目标S3桶的s3:PutObject权限。
  2. 修改output_fn函数,先将图像写入内存字节流,再上传到S3:
import boto3
from io import BytesIO

def output_fn(prediction, content_type):
    assert content_type == 'application/json'
    mask = prediction['mask']
    mask = np.where(mask==2, 255, mask)
    mask = np.where(mask==1, 128, mask)
    trimap = mask.astype(np.uint8)
    trimap_image = Image.fromarray(trimap).resize((prediction['w'], prediction['h'])).convert('L')
    
    # 初始化S3客户端
    s3 = boto3.client('s3')
    bucket_name = 'your-bucket-name'  # 替换为你的S3桶名
    s3_key = 'path/to/test.png'  # 替换为桶内的目标路径
    
    # 将图像写入内存字节流
    img_byte_arr = BytesIO()
    trimap_image.save(img_byte_arr, format='PNG')
    img_byte_arr.seek(0)
    
    # 上传至S3
    s3.upload_fileobj(img_byte_arr, bucket_name, s3_key)
    
    return json.dumps([f'Image saved to s3://{bucket_name}/{s3_key}!'])

如果要通过prediction['save_path']指定S3路径,可解析路径中的桶名和键:

# 假设prediction['save_path']格式为's3://your-bucket/path/to/test.png'
s3_path = prediction['save_path']
bucket_name = s3_path.split('/')[2]
s3_key = '/'.join(s3_path.split('/')[3:])
# 后续上传逻辑同上

二、直接返回PNG图像或Numpy数组

1. 直接返回PNG图像

需要修改output_fn的返回类型和对应content_type,同时调整部署时的序列化/反序列化配置:

  • 修改inference.py中的output_fn:
from io import BytesIO

def output_fn(prediction, content_type):
    mask = prediction['mask']
    mask = np.where(mask==2, 255, mask)
    mask = np.where(mask==1, 128, mask)
    trimap = mask.astype(np.uint8)
    trimap_image = Image.fromarray(trimap).resize((prediction['w'], prediction['h'])).convert('L')
    
    img_byte_arr = BytesIO()
    trimap_image.save(img_byte_arr, format='PNG')
    img_byte_arr.seek(0)
    
    # 返回字节流,并指定content_type为image/png
    return img_byte_arr.getvalue(), 'image/png'
  • 修改部署代码,更换反序列化器:
from sagemaker.deserializers import BytesDeserializer

predictor = model.deploy(
    initial_instance_count=1,
    instance_type=instance_type,
    serializer=JSONSerializer(),
    deserializer=BytesDeserializer(content_type='image/png'),   
)

调用预测后,可直接将返回的字节流保存为PNG文件:

response = predictor.predict(input_data)
with open('output.png', 'wb') as f:
    f.write(response)

2. 直接返回Numpy数组

有两种常用方式:

  • 方案1:返回JSON格式的数组列表
    修改output_fn:
def output_fn(prediction, content_type):
    assert content_type == 'application/json'
    mask = prediction['mask']
    mask = np.where(mask==2, 255, mask)
    mask = np.where(mask==1, 128, mask)
    trimap = mask.astype(np.uint8)
    # 调整尺寸后转为列表
    trimap_resized = np.array(Image.fromarray(trimap).resize((prediction['w'], prediction['h'])))
    return json.dumps(trimap_resized.tolist())

调用后解析JSON还原数组:

response = predictor.predict(input_data)
import numpy as np
trimap_arr = np.array(response)
  • 方案2:返回base64编码的数组(适合大数据量)
import base64

def output_fn(prediction, content_type):
    assert content_type == 'application/json'
    mask = prediction['mask']
    mask = np.where(mask==2, 255, mask)
    mask = np.where(mask==1, 128, mask)
    trimap = mask.astype(np.uint8)
    trimap_resized = np.array(Image.fromarray(trimap).resize((prediction['w'], prediction['h'])))
    # 数组转字节后base64编码
    arr_bytes = trimap_resized.tobytes()
    arr_b64 = base64.b64encode(arr_bytes).decode('utf-8')
    return json.dumps({
        'data': arr_b64,
        'shape': trimap_resized.shape
    })

还原时:

import base64
import numpy as np

response = predictor.predict(input_data)
arr_b64 = response['data']
shape = response['shape']
arr_bytes = base64.b64decode(arr_b64)
trimap_arr = np.frombuffer(arr_bytes, dtype=np.uint8).reshape(shape)

内容的提问来源于stack exchange,提问作者Diego Rodea

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.14 17:42:15