(
self,
device: torch.device,
models: Sequence[nn.Module],
diffusions: Sequence[GaussianDiffusion],
num_points: Sequence[int],
aux_channels: Sequence[str],
model_kwargs_key_filter: Sequence[str] = ("*",),
guidance_scale: Sequence[float] = (3.0, 3.0),
clip_denoised: bool = True,
use_karras: Sequence[bool] = (True, True),
karras_steps: Sequence[int] = (64, 64),
sigma_min: Sequence[float] = (1e-3, 1e-3),
sigma_max: Sequence[float] = (120, 160),
s_churn: Sequence[float] = (3, 0),
)
| 24 | """ |
| 25 | |
| 26 | def __init__( |
| 27 | self, |
| 28 | device: torch.device, |
| 29 | models: Sequence[nn.Module], |
| 30 | diffusions: Sequence[GaussianDiffusion], |
| 31 | num_points: Sequence[int], |
| 32 | aux_channels: Sequence[str], |
| 33 | model_kwargs_key_filter: Sequence[str] = ("*",), |
| 34 | guidance_scale: Sequence[float] = (3.0, 3.0), |
| 35 | clip_denoised: bool = True, |
| 36 | use_karras: Sequence[bool] = (True, True), |
| 37 | karras_steps: Sequence[int] = (64, 64), |
| 38 | sigma_min: Sequence[float] = (1e-3, 1e-3), |
| 39 | sigma_max: Sequence[float] = (120, 160), |
| 40 | s_churn: Sequence[float] = (3, 0), |
| 41 | ): |
| 42 | n = len(models) |
| 43 | assert n > 0 |
| 44 | |
| 45 | if n > 1: |
| 46 | if len(guidance_scale) == 1: |
| 47 | # Don't guide the upsamplers by default. |
| 48 | guidance_scale = list(guidance_scale) + [1.0] * (n - 1) |
| 49 | if len(use_karras) == 1: |
| 50 | use_karras = use_karras * n |
| 51 | if len(karras_steps) == 1: |
| 52 | karras_steps = karras_steps * n |
| 53 | if len(sigma_min) == 1: |
| 54 | sigma_min = sigma_min * n |
| 55 | if len(sigma_max) == 1: |
| 56 | sigma_max = sigma_max * n |
| 57 | if len(s_churn) == 1: |
| 58 | s_churn = s_churn * n |
| 59 | if len(model_kwargs_key_filter) == 1: |
| 60 | model_kwargs_key_filter = model_kwargs_key_filter * n |
| 61 | if len(model_kwargs_key_filter) == 0: |
| 62 | model_kwargs_key_filter = ["*"] * n |
| 63 | assert len(guidance_scale) == n |
| 64 | assert len(use_karras) == n |
| 65 | assert len(karras_steps) == n |
| 66 | assert len(sigma_min) == n |
| 67 | assert len(sigma_max) == n |
| 68 | assert len(s_churn) == n |
| 69 | assert len(model_kwargs_key_filter) == n |
| 70 | |
| 71 | self.device = device |
| 72 | self.num_points = num_points |
| 73 | self.aux_channels = aux_channels |
| 74 | self.model_kwargs_key_filter = model_kwargs_key_filter |
| 75 | self.guidance_scale = guidance_scale |
| 76 | self.clip_denoised = clip_denoised |
| 77 | self.use_karras = use_karras |
| 78 | self.karras_steps = karras_steps |
| 79 | self.sigma_min = sigma_min |
| 80 | self.sigma_max = sigma_max |
| 81 | self.s_churn = s_churn |
| 82 | |
| 83 | self.models = models |
nothing calls this directly
no outgoing calls
no test coverage detected