ADD 增加上传图片返回url接口

This commit is contained in:
2025-07-16 18:11:55 +08:00
parent 5fee351bbd
commit bbea095da9
3 changed files with 118 additions and 4 deletions

View File

@@ -1,5 +1,6 @@
import base64
import io
import tempfile
import torch
from PIL import Image
@@ -71,3 +72,63 @@ def tensor_to_image_bytes(tensor: torch.Tensor, format: str = 'PNG') -> bytes:
buffer.seek(0) # 重置指针到开始位置
return buffer.getvalue()
def tensor_to_tempfile(tensor: torch.Tensor, format: str = 'PNG',
normalize: bool = True, range=None) -> tempfile.NamedTemporaryFile:
"""
将PyTorch张量转换为图像并保存到临时文件
参数:
tensor: 输入的PyTorch张量可以是4D(BCHW)、3D(CHW)或2D(HW)
format: 图像格式,如'PNG''JPEG'
normalize: 是否对张量进行归一化处理
range: 归一化范围,元组(min, max),默认为张量的最小值和最大值
返回:
临时文件对象,关闭后会自动删除
"""
# 处理4D张量 (BCHW),只取第一个样本
if tensor.dim() == 4:
if tensor.size(0) > 1:
print(f"警告: 输入张量包含多个样本,仅使用第一个样本 ({tensor.size(0)} -> 1)")
tensor = tensor[0]
# 确保张量在CPU上
tensor = tensor.cpu()
tensor = tensor.permute(2,0,1)
# 归一化处理
if normalize:
if range is None:
min_val, max_val = tensor.min(), tensor.max()
else:
min_val, max_val = range
if max_val > min_val:
tensor = (tensor - min_val) / (max_val - min_val)
else:
tensor = torch.zeros_like(tensor)
# 转换为PIL图像
if tensor.dim() == 2: # HW格式 (灰度图)
pil_img = transforms.ToPILImage()(tensor.unsqueeze(0)) # 添加通道维度
elif tensor.dim() == 3: # CHW格式
pil_img = transforms.ToPILImage()(tensor)
else:
raise ValueError(f"不支持的张量维度: {tensor.dim()}")
# 创建临时文件
temp_file = tempfile.NamedTemporaryFile(suffix=f'.{format.lower()}', delete=False)
try:
# 保存图像到临时文件
pil_img.save(temp_file, format=format)
except Exception as e:
# 发生错误时删除临时文件
temp_file.close()
raise e
temp_file.close() # 关闭文件但不删除
return temp_file