(
self, batch_size: int, model_kwargs: Dict[str, Any]
)
| 94 | return samples |
| 95 | |
| 96 | def sample_batch_progressive( |
| 97 | self, batch_size: int, model_kwargs: Dict[str, Any] |
| 98 | ) -> Iterator[torch.Tensor]: |
| 99 | samples = None |
| 100 | for ( |
| 101 | model, |
| 102 | diffusion, |
| 103 | stage_num_points, |
| 104 | stage_guidance_scale, |
| 105 | stage_use_karras, |
| 106 | stage_karras_steps, |
| 107 | stage_sigma_min, |
| 108 | stage_sigma_max, |
| 109 | stage_s_churn, |
| 110 | stage_key_filter, |
| 111 | ) in zip( |
| 112 | self.models, |
| 113 | self.diffusions, |
| 114 | self.num_points, |
| 115 | self.guidance_scale, |
| 116 | self.use_karras, |
| 117 | self.karras_steps, |
| 118 | self.sigma_min, |
| 119 | self.sigma_max, |
| 120 | self.s_churn, |
| 121 | self.model_kwargs_key_filter, |
| 122 | ): |
| 123 | stage_model_kwargs = model_kwargs.copy() |
| 124 | if stage_key_filter != "*": |
| 125 | use_keys = set(stage_key_filter.split(",")) |
| 126 | stage_model_kwargs = {k: v for k, v in stage_model_kwargs.items() if k in use_keys} |
| 127 | if samples is not None: |
| 128 | stage_model_kwargs["low_res"] = samples |
| 129 | if hasattr(model, "cached_model_kwargs"): |
| 130 | stage_model_kwargs = model.cached_model_kwargs(batch_size, stage_model_kwargs) |
| 131 | sample_shape = (batch_size, 3 + len(self.aux_channels), stage_num_points) |
| 132 | |
| 133 | if stage_guidance_scale != 1 and stage_guidance_scale != 0: |
| 134 | for k, v in stage_model_kwargs.copy().items(): |
| 135 | stage_model_kwargs[k] = torch.cat([v, torch.zeros_like(v)], dim=0) |
| 136 | |
| 137 | if stage_use_karras: |
| 138 | samples_it = karras_sample_progressive( |
| 139 | diffusion=diffusion, |
| 140 | model=model, |
| 141 | shape=sample_shape, |
| 142 | steps=stage_karras_steps, |
| 143 | clip_denoised=self.clip_denoised, |
| 144 | model_kwargs=stage_model_kwargs, |
| 145 | device=self.device, |
| 146 | sigma_min=stage_sigma_min, |
| 147 | sigma_max=stage_sigma_max, |
| 148 | s_churn=stage_s_churn, |
| 149 | guidance_scale=stage_guidance_scale, |
| 150 | ) |
| 151 | else: |
| 152 | internal_batch_size = batch_size |
| 153 | if stage_guidance_scale: |
no test coverage detected