Add Stable Diffusion 1.5 support with data-parallel inference
Add base15.yml for SD 1.5 (PyTorch weights via from_pt, PNDM/epsilon scheduler) and wire generate.py to it: - Build the sampler from the checkpoint's scheduler config via create_scheduler instead of a hardcoded DDIM scheduler, and iterate the full PNDM schedule (skip_prk_steps emits one extra timestep). - Shard the latent batch over the data axis with sharding constraints plus out_shardings so inference runs data parallel instead of replicating the whole batch on every device. Sub-device batches replicate. - Make override_scheduler_config tolerant of scheduler configs that omit keys (e.g. SD 1.5's PNDM config).
C
csgoogle committed
164ca87b0f66e6d42bc02dc44a3a145ed7e269b4
Parent: b2d31df