如何更高效运行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
相关产品推荐
相关产品推荐

