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

如何更高效运行pil_to_latents函数?合并多轮列表推导为单次循环

解决方案

你可以把所有图片处理逻辑合并到单次循环中,避免多次遍历数据集列表,同时保持代码清晰。直接对每张图片依次执行转换、张量处理、设备迁移和编码操作即可:

def pil_to_latents(dataset):
    '''
    Function to convert image to latents
    '''
    latents = []
    for img in dataset['train']['image']:
        # 单张图片的完整处理流程
        rgb_img = img.convert('RGB').resize((config.image_size, config.image_size))
        tensor_img = tfms.ToTensor()(rgb_img).unsqueeze(0) * 2.0 - 1.0
        tensor_img = tensor_img.to(device="cuda", dtype=torch.float16)
        latent = vae.encode(tensor_img).latent_dist.sample() * 0.18215
        latents.append(latent)
    
    dataset['train']['latents'] = latents
    return dataset

补充说明

  • 单次遍历即可完成所有操作,减少了原代码多次列表推导带来的重复遍历开销
  • 如果需要复用这套处理逻辑,可以自己封装一个简单的可调用类,绕过torchvision.Compose的局限:
class ImageToLatentTransform:
    def __init__(self, image_size, device="cuda"):
        self.image_size = image_size
        self.device = device
    
    def __call__(self, img):
        rgb_img = img.convert('RGB').resize((self.image_size, self.image_size))
        tensor_img = tfms.ToTensor()(rgb_img).unsqueeze(0) * 2.0 - 1.0
        tensor_img = tensor_img.to(device=self.device, dtype=torch.float16)
        latent = vae.encode(tensor_img).latent_dist.sample() * 0.18215
        return latent

# 使用示例
transform = ImageToLatentTransform(config.image_size)
dataset['train']['latents'] = [transform(img) for img in dataset['train']['image']]

这个自定义类可以像Compose一样复用逻辑,同时完全支持你的自定义操作,灵活性更强。

内容的提问来源于stack exchange,提问作者Nicolas Pereyra Zorraquin

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.24 19:25:09