Team Ai
Apppublic

hugging-apps/echo-memory

sourceHugging Faceupdated 3mo agoView on Hugging Face
0likes
cog_vae.py519 linesDownload Raw Back to root
1import torch2from einops import rearrange, repeat3from .tiler import TileWorker2Dto3D4 5 6 7class Downsample3D(torch.nn.Module):8    def __init__(9        self,10        in_channels: int,11        out_channels: int,12        kernel_size: int = 3,13        stride: int = 2,14        padding: int = 0,15        compress_time: bool = False,16    ):17        super().__init__()18 19        self.conv = torch.nn.Conv2d(in_channels, out_channels, kernel_size=kernel_size, stride=stride, padding=padding)20        self.compress_time = compress_time21 22    def forward(self, x: torch.Tensor, xq: torch.Tensor) -> torch.Tensor:23        if self.compress_time:24            batch_size, channels, frames, height, width = x.shape25 26            # (batch_size, channels, frames, height, width) -> (batch_size, height, width, channels, frames) -> (batch_size * height * width, channels, frames)27            x = x.permute(0, 3, 4, 1, 2).reshape(batch_size * height * width, channels, frames)28 29            if x.shape[-1] % 2 == 1:30                x_first, x_rest = x[..., 0], x[..., 1:]31                if x_rest.shape[-1] > 0:32                    # (batch_size * height * width, channels, frames - 1) -> (batch_size * height * width, channels, (frames - 1) // 2)33                    x_rest = torch.nn.functional.avg_pool1d(x_rest, kernel_size=2, stride=2)34 35                x = torch.cat([x_first[..., None], x_rest], dim=-1)36                # (batch_size * height * width, channels, (frames // 2) + 1) -> (batch_size, height, width, channels, (frames // 2) + 1) -> (batch_size, channels, (frames // 2) + 1, height, width)37                x = x.reshape(batch_size, height, width, channels, x.shape[-1]).permute(0, 3, 4, 1, 2)38            else:39                # (batch_size * height * width, channels, frames) -> (batch_size * height * width, channels, frames // 2)40                x = torch.nn.functional.avg_pool1d(x, kernel_size=2, stride=2)41                # (batch_size * height * width, channels, frames // 2) -> (batch_size, height, width, channels, frames // 2) -> (batch_size, channels, frames // 2, height, width)42                x = x.reshape(batch_size, height, width, channels, x.shape[-1]).permute(0, 3, 4, 1, 2)43 44        # Pad the tensor45        pad = (0, 1, 0, 1)46        x = torch.nn.functional.pad(x, pad, mode="constant", value=0)47        batch_size, channels, frames, height, width = x.shape48        # (batch_size, channels, frames, height, width) -> (batch_size, frames, channels, height, width) -> (batch_size * frames, channels, height, width)49        x = x.permute(0, 2, 1, 3, 4).reshape(batch_size * frames, channels, height, width)50        x = self.conv(x)51        # (batch_size * frames, channels, height, width) -> (batch_size, frames, channels, height, width) -> (batch_size, channels, frames, height, width)52        x = x.reshape(batch_size, frames, x.shape[1], x.shape[2], x.shape[3]).permute(0, 2, 1, 3, 4)53        return x54 55 56 57class Upsample3D(torch.nn.Module):58    def __init__(59        self,60        in_channels: int,61        out_channels: int,62        kernel_size: int = 3,63        stride: int = 1,64        padding: int = 1,65        compress_time: bool = False,66    ) -> None:67        super().__init__()68        self.conv = torch.nn.Conv2d(in_channels, out_channels, kernel_size=kernel_size, stride=stride, padding=padding)69        self.compress_time = compress_time70 71    def forward(self, inputs: torch.Tensor, xq: torch.Tensor) -> torch.Tensor:72        if self.compress_time:73            if inputs.shape[2] > 1 and inputs.shape[2] % 2 == 1:74                # split first frame75                x_first, x_rest = inputs[:, :, 0], inputs[:, :, 1:]76 77                x_first = torch.nn.functional.interpolate(x_first, scale_factor=2.0)78                x_rest = torch.nn.functional.interpolate(x_rest, scale_factor=2.0)79                x_first = x_first[:, :, None, :, :]80                inputs = torch.cat([x_first, x_rest], dim=2)81            elif inputs.shape[2] > 1:82                inputs = torch.nn.functional.interpolate(inputs, scale_factor=2.0)83            else:84                inputs = inputs.squeeze(2)85                inputs = torch.nn.functional.interpolate(inputs, scale_factor=2.0)86                inputs = inputs[:, :, None, :, :]87        else:88            # only interpolate 2D89            b, c, t, h, w = inputs.shape90            inputs = inputs.permute(0, 2, 1, 3, 4).reshape(b * t, c, h, w)91            inputs = torch.nn.functional.interpolate(inputs, scale_factor=2.0)92            inputs = inputs.reshape(b, t, c, *inputs.shape[2:]).permute(0, 2, 1, 3, 4)93 94        b, c, t, h, w = inputs.shape95        inputs = inputs.permute(0, 2, 1, 3, 4).reshape(b * t, c, h, w)96        inputs = self.conv(inputs)97        inputs = inputs.reshape(b, t, *inputs.shape[1:]).permute(0, 2, 1, 3, 4)98 99        return inputs100 101 102 103class CogVideoXSpatialNorm3D(torch.nn.Module):104    def __init__(self, f_channels, zq_channels, groups):105        super().__init__()106        self.norm_layer = torch.nn.GroupNorm(num_channels=f_channels, num_groups=groups, eps=1e-6, affine=True)107        self.conv_y = torch.nn.Conv3d(zq_channels, f_channels, kernel_size=1, stride=1)108        self.conv_b = torch.nn.Conv3d(zq_channels, f_channels, kernel_size=1, stride=1)109 110 111    def forward(self, f: torch.Tensor, zq: torch.Tensor) -> torch.Tensor:112        if f.shape[2] > 1 and f.shape[2] % 2 == 1:113            f_first, f_rest = f[:, :, :1], f[:, :, 1:]114            f_first_size, f_rest_size = f_first.shape[-3:], f_rest.shape[-3:]115            z_first, z_rest = zq[:, :, :1], zq[:, :, 1:]116            z_first = torch.nn.functional.interpolate(z_first, size=f_first_size)117            z_rest = torch.nn.functional.interpolate(z_rest, size=f_rest_size)118            zq = torch.cat([z_first, z_rest], dim=2)119        else:120            zq = torch.nn.functional.interpolate(zq, size=f.shape[-3:])121 122        norm_f = self.norm_layer(f)123        new_f = norm_f * self.conv_y(zq) + self.conv_b(zq)124        return new_f125 126 127 128class Resnet3DBlock(torch.nn.Module):129    def __init__(self, in_channels, out_channels, spatial_norm_dim, groups, eps=1e-6, use_conv_shortcut=False):130        super().__init__()131        self.nonlinearity = torch.nn.SiLU()132        if spatial_norm_dim is None:133            self.norm1 = torch.nn.GroupNorm(num_channels=in_channels, num_groups=groups, eps=eps)134            self.norm2 = torch.nn.GroupNorm(num_channels=out_channels, num_groups=groups, eps=eps)135        else:136            self.norm1 = CogVideoXSpatialNorm3D(in_channels, spatial_norm_dim, groups)137            self.norm2 = CogVideoXSpatialNorm3D(out_channels, spatial_norm_dim, groups)138 139        self.conv1 = CachedConv3d(in_channels, out_channels, kernel_size=3, padding=(0, 1, 1))140 141        self.conv2 = CachedConv3d(out_channels, out_channels, kernel_size=3, padding=(0, 1, 1))142 143        if in_channels != out_channels:144            if use_conv_shortcut:145                self.conv_shortcut = CachedConv3d(in_channels, out_channels, kernel_size=3, padding=(0, 1, 1))146            else:147                self.conv_shortcut = torch.nn.Conv3d(in_channels, out_channels, kernel_size=1)148        else:149            self.conv_shortcut = lambda x: x150 151 152    def forward(self, hidden_states, zq):153        residual = hidden_states154 155        hidden_states = self.norm1(hidden_states, zq) if isinstance(self.norm1, CogVideoXSpatialNorm3D) else self.norm1(hidden_states)156        hidden_states = self.nonlinearity(hidden_states)157        hidden_states = self.conv1(hidden_states)158 159        hidden_states = self.norm2(hidden_states, zq) if isinstance(self.norm2, CogVideoXSpatialNorm3D) else self.norm2(hidden_states)160        hidden_states = self.nonlinearity(hidden_states)161        hidden_states = self.conv2(hidden_states)162 163        hidden_states = hidden_states + self.conv_shortcut(residual)164 165        return hidden_states166    167 168 169class CachedConv3d(torch.nn.Conv3d):170    def __init__(self, in_channels, out_channels, kernel_size, stride=1, padding=0):171        super().__init__(in_channels, out_channels, kernel_size=kernel_size, stride=stride, padding=padding)172        self.cached_tensor = None173 174 175    def clear_cache(self):176        self.cached_tensor = None177    178 179    def forward(self, input: torch.Tensor, use_cache = True) -> torch.Tensor:180        if use_cache:181            if self.cached_tensor is None:182                self.cached_tensor = torch.concat([input[:, :, :1]] * 2, dim=2)183            input = torch.concat([self.cached_tensor, input], dim=2)184            self.cached_tensor = input[:, :, -2:]185        return super().forward(input)186 187 188 189class CogVAEDecoder(torch.nn.Module):190    def __init__(self):191        super().__init__()192        self.scaling_factor = 0.7193        self.conv_in = CachedConv3d(16, 512, kernel_size=3, stride=1, padding=(0, 1, 1))194 195        self.blocks = torch.nn.ModuleList([196            Resnet3DBlock(512, 512, 16, 32),197            Resnet3DBlock(512, 512, 16, 32),198            Resnet3DBlock(512, 512, 16, 32),199            Resnet3DBlock(512, 512, 16, 32),200            Resnet3DBlock(512, 512, 16, 32),201            Resnet3DBlock(512, 512, 16, 32),202            Upsample3D(512, 512, compress_time=True),203            Resnet3DBlock(512, 256, 16, 32),204            Resnet3DBlock(256, 256, 16, 32),205            Resnet3DBlock(256, 256, 16, 32),206            Resnet3DBlock(256, 256, 16, 32),207            Upsample3D(256, 256, compress_time=True),208            Resnet3DBlock(256, 256, 16, 32),209            Resnet3DBlock(256, 256, 16, 32),210            Resnet3DBlock(256, 256, 16, 32),211            Resnet3DBlock(256, 256, 16, 32),212            Upsample3D(256, 256, compress_time=False),213            Resnet3DBlock(256, 128, 16, 32),214            Resnet3DBlock(128, 128, 16, 32),215            Resnet3DBlock(128, 128, 16, 32),216            Resnet3DBlock(128, 128, 16, 32),217        ])218 219        self.norm_out = CogVideoXSpatialNorm3D(128, 16, 32)220        self.conv_act = torch.nn.SiLU()221        self.conv_out = CachedConv3d(128, 3, kernel_size=3, stride=1, padding=(0, 1, 1))222 223 224    def forward(self, sample):225        sample = sample / self.scaling_factor226        hidden_states = self.conv_in(sample)227 228        for block in self.blocks:229            hidden_states = block(hidden_states, sample)230        231        hidden_states = self.norm_out(hidden_states, sample)232        hidden_states = self.conv_act(hidden_states)233        hidden_states = self.conv_out(hidden_states)234 235        return hidden_states236    237 238    def decode_video(self, sample, tiled=True, tile_size=(60, 90), tile_stride=(30, 45), progress_bar=lambda x:x):239        if tiled:240            B, C, T, H, W = sample.shape241            return TileWorker2Dto3D().tiled_forward(242                forward_fn=lambda x: self.decode_small_video(x),243                model_input=sample,244                tile_size=tile_size, tile_stride=tile_stride,245                tile_device=sample.device, tile_dtype=sample.dtype,246                computation_device=sample.device, computation_dtype=sample.dtype,247                scales=(3/16, (T//2*8+T%2)/T, 8, 8),248                progress_bar=progress_bar249            )250        else:251            return self.decode_small_video(sample)252    253 254    def decode_small_video(self, sample):255        B, C, T, H, W = sample.shape256        computation_device = self.conv_in.weight.device257        computation_dtype = self.conv_in.weight.dtype258        value = []259        for i in range(T//2):260            tl = i*2 + T%2 - (T%2 and i==0)261            tr = i*2 + 2 + T%2262            model_input = sample[:, :, tl: tr, :, :].to(dtype=computation_dtype, device=computation_device)263            model_output = self.forward(model_input).to(dtype=sample.dtype, device=sample.device)264            value.append(model_output)265        value = torch.concat(value, dim=2)266        for name, module in self.named_modules():267            if isinstance(module, CachedConv3d):268                module.clear_cache()269        return value270    271 272    @staticmethod273    def state_dict_converter():274        return CogVAEDecoderStateDictConverter()275    276 277 278class CogVAEEncoder(torch.nn.Module):279    def __init__(self):280        super().__init__()281        self.scaling_factor = 0.7282        self.conv_in = CachedConv3d(3, 128, kernel_size=3, stride=1, padding=(0, 1, 1))283 284        self.blocks = torch.nn.ModuleList([285            Resnet3DBlock(128, 128, None, 32),286            Resnet3DBlock(128, 128, None, 32),287            Resnet3DBlock(128, 128, None, 32),288            Downsample3D(128, 128, compress_time=True),289            Resnet3DBlock(128, 256, None, 32),290            Resnet3DBlock(256, 256, None, 32),291            Resnet3DBlock(256, 256, None, 32),292            Downsample3D(256, 256, compress_time=True),293            Resnet3DBlock(256, 256, None, 32),294            Resnet3DBlock(256, 256, None, 32),295            Resnet3DBlock(256, 256, None, 32),296            Downsample3D(256, 256, compress_time=False),297            Resnet3DBlock(256, 512, None, 32),298            Resnet3DBlock(512, 512, None, 32),299            Resnet3DBlock(512, 512, None, 32),300            Resnet3DBlock(512, 512, None, 32),301            Resnet3DBlock(512, 512, None, 32),302        ])303 304        self.norm_out = torch.nn.GroupNorm(32, 512, eps=1e-06, affine=True)305        self.conv_act = torch.nn.SiLU()306        self.conv_out = CachedConv3d(512, 32, kernel_size=3, stride=1, padding=(0, 1, 1))307 308 309    def forward(self, sample):310        hidden_states = self.conv_in(sample)311 312        for block in self.blocks:313            hidden_states = block(hidden_states, sample)314        315        hidden_states = self.norm_out(hidden_states)316        hidden_states = self.conv_act(hidden_states)317        hidden_states = self.conv_out(hidden_states)[:, :16]318        hidden_states = hidden_states * self.scaling_factor319 320        return hidden_states321    322 323    def encode_video(self, sample, tiled=True, tile_size=(60, 90), tile_stride=(30, 45), progress_bar=lambda x:x):324        if tiled:325            B, C, T, H, W = sample.shape326            return TileWorker2Dto3D().tiled_forward(327                forward_fn=lambda x: self.encode_small_video(x),328                model_input=sample,329                tile_size=(i * 8 for i in tile_size), tile_stride=(i * 8 for i in tile_stride),330                tile_device=sample.device, tile_dtype=sample.dtype,331                computation_device=sample.device, computation_dtype=sample.dtype,332                scales=(16/3, (T//4+T%2)/T, 1/8, 1/8),333                progress_bar=progress_bar334            )335        else:336            return self.encode_small_video(sample)337    338 339    def encode_small_video(self, sample):340        B, C, T, H, W = sample.shape341        computation_device = self.conv_in.weight.device342        computation_dtype = self.conv_in.weight.dtype343        value = []344        for i in range(T//8):345            t = i*8 + T%2 - (T%2 and i==0)346            t_ = i*8 + 8 + T%2347            model_input = sample[:, :, t: t_, :, :].to(dtype=computation_dtype, device=computation_device)348            model_output = self.forward(model_input).to(dtype=sample.dtype, device=sample.device)349            value.append(model_output)350        value = torch.concat(value, dim=2)351        for name, module in self.named_modules():352            if isinstance(module, CachedConv3d):353                module.clear_cache()354        return value355    356 357    @staticmethod358    def state_dict_converter():359        return CogVAEEncoderStateDictConverter()360 361 362 363class CogVAEEncoderStateDictConverter:364    def __init__(self):365        pass366 367 368    def from_diffusers(self, state_dict):369        rename_dict = {370            "encoder.conv_in.conv.weight": "conv_in.weight",371            "encoder.conv_in.conv.bias": "conv_in.bias",372            "encoder.down_blocks.0.downsamplers.0.conv.weight": "blocks.3.conv.weight",373            "encoder.down_blocks.0.downsamplers.0.conv.bias": "blocks.3.conv.bias",374            "encoder.down_blocks.1.downsamplers.0.conv.weight": "blocks.7.conv.weight",375            "encoder.down_blocks.1.downsamplers.0.conv.bias": "blocks.7.conv.bias",376            "encoder.down_blocks.2.downsamplers.0.conv.weight": "blocks.11.conv.weight",377            "encoder.down_blocks.2.downsamplers.0.conv.bias": "blocks.11.conv.bias",378            "encoder.norm_out.weight": "norm_out.weight",379            "encoder.norm_out.bias": "norm_out.bias",380            "encoder.conv_out.conv.weight": "conv_out.weight",381            "encoder.conv_out.conv.bias": "conv_out.bias",382        }383        prefix_dict = {384            "encoder.down_blocks.0.resnets.0.": "blocks.0.",385            "encoder.down_blocks.0.resnets.1.": "blocks.1.",386            "encoder.down_blocks.0.resnets.2.": "blocks.2.",387            "encoder.down_blocks.1.resnets.0.": "blocks.4.",388            "encoder.down_blocks.1.resnets.1.": "blocks.5.",389            "encoder.down_blocks.1.resnets.2.": "blocks.6.",390            "encoder.down_blocks.2.resnets.0.": "blocks.8.",391            "encoder.down_blocks.2.resnets.1.": "blocks.9.",392            "encoder.down_blocks.2.resnets.2.": "blocks.10.",393            "encoder.down_blocks.3.resnets.0.": "blocks.12.",394            "encoder.down_blocks.3.resnets.1.": "blocks.13.",395            "encoder.down_blocks.3.resnets.2.": "blocks.14.",396            "encoder.mid_block.resnets.0.": "blocks.15.",397            "encoder.mid_block.resnets.1.": "blocks.16.",398        }399        suffix_dict = {400            "norm1.norm_layer.weight": "norm1.norm_layer.weight",401            "norm1.norm_layer.bias": "norm1.norm_layer.bias",402            "norm1.conv_y.conv.weight": "norm1.conv_y.weight",403            "norm1.conv_y.conv.bias": "norm1.conv_y.bias",404            "norm1.conv_b.conv.weight": "norm1.conv_b.weight",405            "norm1.conv_b.conv.bias": "norm1.conv_b.bias",406            "norm2.norm_layer.weight": "norm2.norm_layer.weight",407            "norm2.norm_layer.bias": "norm2.norm_layer.bias",408            "norm2.conv_y.conv.weight": "norm2.conv_y.weight",409            "norm2.conv_y.conv.bias": "norm2.conv_y.bias",410            "norm2.conv_b.conv.weight": "norm2.conv_b.weight",411            "norm2.conv_b.conv.bias": "norm2.conv_b.bias",412            "conv1.conv.weight": "conv1.weight",413            "conv1.conv.bias": "conv1.bias",414            "conv2.conv.weight": "conv2.weight",415            "conv2.conv.bias": "conv2.bias",416            "conv_shortcut.weight": "conv_shortcut.weight",417            "conv_shortcut.bias": "conv_shortcut.bias",418            "norm1.weight": "norm1.weight",419            "norm1.bias": "norm1.bias",420            "norm2.weight": "norm2.weight",421            "norm2.bias": "norm2.bias",422        }423        state_dict_ = {}424        for name, param in state_dict.items():425            if name in rename_dict:426                state_dict_[rename_dict[name]] = param427            else:428                for prefix in prefix_dict:429                    if name.startswith(prefix):430                        suffix = name[len(prefix):]431                        state_dict_[prefix_dict[prefix] + suffix_dict[suffix]] = param432        return state_dict_433    434 435    def from_civitai(self, state_dict):436        return self.from_diffusers(state_dict)437 438 439 440class CogVAEDecoderStateDictConverter:441    def __init__(self):442        pass443 444 445    def from_diffusers(self, state_dict):446        rename_dict = {447            "decoder.conv_in.conv.weight": "conv_in.weight",448            "decoder.conv_in.conv.bias": "conv_in.bias",449            "decoder.up_blocks.0.upsamplers.0.conv.weight": "blocks.6.conv.weight",450            "decoder.up_blocks.0.upsamplers.0.conv.bias": "blocks.6.conv.bias",451            "decoder.up_blocks.1.upsamplers.0.conv.weight": "blocks.11.conv.weight",452            "decoder.up_blocks.1.upsamplers.0.conv.bias": "blocks.11.conv.bias",453            "decoder.up_blocks.2.upsamplers.0.conv.weight": "blocks.16.conv.weight",454            "decoder.up_blocks.2.upsamplers.0.conv.bias": "blocks.16.conv.bias",455            "decoder.norm_out.norm_layer.weight": "norm_out.norm_layer.weight",456            "decoder.norm_out.norm_layer.bias": "norm_out.norm_layer.bias",457            "decoder.norm_out.conv_y.conv.weight": "norm_out.conv_y.weight",458            "decoder.norm_out.conv_y.conv.bias": "norm_out.conv_y.bias",459            "decoder.norm_out.conv_b.conv.weight": "norm_out.conv_b.weight",460            "decoder.norm_out.conv_b.conv.bias": "norm_out.conv_b.bias",461            "decoder.conv_out.conv.weight": "conv_out.weight",462            "decoder.conv_out.conv.bias": "conv_out.bias"463        }464        prefix_dict = {465            "decoder.mid_block.resnets.0.": "blocks.0.",466            "decoder.mid_block.resnets.1.": "blocks.1.",467            "decoder.up_blocks.0.resnets.0.": "blocks.2.",468            "decoder.up_blocks.0.resnets.1.": "blocks.3.",469            "decoder.up_blocks.0.resnets.2.": "blocks.4.",470            "decoder.up_blocks.0.resnets.3.": "blocks.5.",471            "decoder.up_blocks.1.resnets.0.": "blocks.7.",472            "decoder.up_blocks.1.resnets.1.": "blocks.8.",473            "decoder.up_blocks.1.resnets.2.": "blocks.9.",474            "decoder.up_blocks.1.resnets.3.": "blocks.10.",475            "decoder.up_blocks.2.resnets.0.": "blocks.12.",476            "decoder.up_blocks.2.resnets.1.": "blocks.13.",477            "decoder.up_blocks.2.resnets.2.": "blocks.14.",478            "decoder.up_blocks.2.resnets.3.": "blocks.15.",479            "decoder.up_blocks.3.resnets.0.": "blocks.17.",480            "decoder.up_blocks.3.resnets.1.": "blocks.18.",481            "decoder.up_blocks.3.resnets.2.": "blocks.19.",482            "decoder.up_blocks.3.resnets.3.": "blocks.20.",483        }484        suffix_dict = {485            "norm1.norm_layer.weight": "norm1.norm_layer.weight",486            "norm1.norm_layer.bias": "norm1.norm_layer.bias",487            "norm1.conv_y.conv.weight": "norm1.conv_y.weight",488            "norm1.conv_y.conv.bias": "norm1.conv_y.bias",489            "norm1.conv_b.conv.weight": "norm1.conv_b.weight",490            "norm1.conv_b.conv.bias": "norm1.conv_b.bias",491            "norm2.norm_layer.weight": "norm2.norm_layer.weight",492            "norm2.norm_layer.bias": "norm2.norm_layer.bias",493            "norm2.conv_y.conv.weight": "norm2.conv_y.weight",494            "norm2.conv_y.conv.bias": "norm2.conv_y.bias",495            "norm2.conv_b.conv.weight": "norm2.conv_b.weight",496            "norm2.conv_b.conv.bias": "norm2.conv_b.bias",497            "conv1.conv.weight": "conv1.weight",498            "conv1.conv.bias": "conv1.bias",499            "conv2.conv.weight": "conv2.weight",500            "conv2.conv.bias": "conv2.bias",501            "conv_shortcut.weight": "conv_shortcut.weight",502            "conv_shortcut.bias": "conv_shortcut.bias",503        }504        state_dict_ = {}505        for name, param in state_dict.items():506            if name in rename_dict:507                state_dict_[rename_dict[name]] = param508            else:509                for prefix in prefix_dict:510                    if name.startswith(prefix):511                        suffix = name[len(prefix):]512                        state_dict_[prefix_dict[prefix] + suffix_dict[suffix]] = param513        return state_dict_514    515 516    def from_civitai(self, state_dict):517        return self.from_diffusers(state_dict)518 519