| 604 | |
| 605 | |
| 606 | class KDiffusionSampler: |
| 607 | def __init__(self, m, sampler): |
| 608 | self.model = m |
| 609 | self.model_wrap = K.external.CompVisDenoiser(m) |
| 610 | self.schedule = sampler |
| 611 | |
| 612 | def get_sampler_name(self): |
| 613 | return self.schedule |
| 614 | |
| 615 | def sample( |
| 616 | self, |
| 617 | S, |
| 618 | conditioning, |
| 619 | batch_size, |
| 620 | shape, |
| 621 | verbose, |
| 622 | unconditional_guidance_scale, |
| 623 | unconditional_conditioning, |
| 624 | eta, |
| 625 | x_T, |
| 626 | img_callback: Callable = None, |
| 627 | ): |
| 628 | sigmas = self.model_wrap.get_sigmas(S) |
| 629 | x = x_T * sigmas[0] |
| 630 | model_wrap_cfg = CFGDenoiser(self.model_wrap) |
| 631 | |
| 632 | samples_ddim = K.sampling.__dict__[f"sample_{self.schedule}"]( |
| 633 | model_wrap_cfg, |
| 634 | x, |
| 635 | sigmas, |
| 636 | extra_args={ |
| 637 | "cond": conditioning, |
| 638 | "uncond": unconditional_conditioning, |
| 639 | "cond_scale": unconditional_guidance_scale, |
| 640 | }, |
| 641 | disable=False, |
| 642 | callback=partial(KDiffusionSampler.img_callback_wrapper, img_callback), |
| 643 | ) |
| 644 | |
| 645 | return samples_ddim, None |
| 646 | |
| 647 | @classmethod |
| 648 | def img_callback_wrapper(cls, callback: Callable, *args): |
| 649 | """Converts a KDiffusion callback to the standard img_callback""" |
| 650 | if callback: |
| 651 | arg_dict = args[0] |
| 652 | callback(image_sample=arg_dict["denoised"], iter_num=arg_dict["i"]) |
| 653 | |
| 654 | |
| 655 | def create_random_tensors(shape, seeds): |
no outgoing calls
no test coverage detected