(self, sigma: torch.Tensor, step_sigmas: torch.Tensor)
| 28 | return (model,) |
| 29 | |
| 30 | def find_step(self, sigma: torch.Tensor, step_sigmas: torch.Tensor): |
| 31 | for i, step_sigma in enumerate(step_sigmas): |
| 32 | if step_sigma <= sigma: |
| 33 | return i |
| 34 | return len(step_sigmas) - 1 |
| 35 | |
| 36 | def forward( |
| 37 | self, sigma: torch.Tensor, denoise_mask: torch.Tensor, extra_options: dict |