SIGN IN SIGN UP

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