| 39 | """ |
| 40 | |
| 41 | def __init__(self, adapters: List["T2IAdapter"]): |
| 42 | super(MultiAdapter, self).__init__() |
| 43 | |
| 44 | self.num_adapter = len(adapters) |
| 45 | self.adapters = nn.ModuleList(adapters) |
| 46 | |
| 47 | if len(adapters) == 0: |
| 48 | raise ValueError("Expecting at least one adapter") |
| 49 | |
| 50 | if len(adapters) == 1: |
| 51 | raise ValueError("For a single adapter, please use the `T2IAdapter` class instead of `MultiAdapter`") |
| 52 | |
| 53 | # The outputs from each adapter are added together with a weight. |
| 54 | # This means that the change in dimensions from downsampling must |
| 55 | # be the same for all adapters. Inductively, it also means the |
| 56 | # downscale_factor and total_downscale_factor must be the same for all |
| 57 | # adapters. |
| 58 | first_adapter_total_downscale_factor = adapters[0].total_downscale_factor |
| 59 | first_adapter_downscale_factor = adapters[0].downscale_factor |
| 60 | for idx in range(1, len(adapters)): |
| 61 | if ( |
| 62 | adapters[idx].total_downscale_factor != first_adapter_total_downscale_factor |
| 63 | or adapters[idx].downscale_factor != first_adapter_downscale_factor |
| 64 | ): |
| 65 | raise ValueError( |
| 66 | f"Expecting all adapters to have the same downscaling behavior, but got:\n" |
| 67 | f"adapters[0].total_downscale_factor={first_adapter_total_downscale_factor}\n" |
| 68 | f"adapters[0].downscale_factor={first_adapter_downscale_factor}\n" |
| 69 | f"adapter[`{idx}`].total_downscale_factor={adapters[idx].total_downscale_factor}\n" |
| 70 | f"adapter[`{idx}`].downscale_factor={adapters[idx].downscale_factor}" |
| 71 | ) |
| 72 | |
| 73 | self.total_downscale_factor = first_adapter_total_downscale_factor |
| 74 | self.downscale_factor = first_adapter_downscale_factor |
| 75 | |
| 76 | def forward(self, xs: torch.Tensor, adapter_weights: Optional[List[float]] = None) -> List[torch.Tensor]: |
| 77 | r""" |