(tensor, vae)
| 113 | |
| 114 | |
| 115 | def vae_encode(tensor, vae): |
| 116 | # tensor values already in range [-1, 1] here |
| 117 | p = next(vae.encoder.parameters()) |
| 118 | # TODO: the official code would call vae.encode_image() when it detects frames=1. |
| 119 | # Should we use the image encoder (separate model)? |
| 120 | return vae.encode(tensor.to(p.device, p.dtype)) |
| 121 | |
| 122 | |
| 123 | def dataset_config_validation(config): |