File size: 1,246 Bytes
c0f018c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
"""Sol-branded adapter for the matching SANA Diffusers autoencoder."""

from __future__ import annotations

from typing import Any

from torch import nn


class AutoencoderDCSol(nn.Module):
    """Thin wrapper around the frozen SANA 1.1 ``AutoencoderDC`` decoder.

    The learned decoder weights stay in the upstream Diffusers repository;
    this adapter gives SolPix a stable, descriptive component name.
    """

    model_id = "mit-han-lab/dc-ae-f32c32-sana-1.1-diffusers"
    revision = "df0d9d634aea77793c1fb685d4b9db092c99e686"

    def __init__(self, autoencoder: nn.Module):
        super().__init__()
        self.autoencoder = autoencoder

    @classmethod
    def from_pretrained(cls, model_id: str | None = None, **kwargs: Any) -> "AutoencoderDCSol":
        from diffusers import AutoencoderDC

        kwargs.setdefault("revision", cls.revision)
        return cls(AutoencoderDC.from_pretrained(model_id or cls.model_id, **kwargs))

    @property
    def config(self) -> Any:
        return self.autoencoder.config

    def decode(self, *args: Any, **kwargs: Any) -> Any:
        return self.autoencoder.decode(*args, **kwargs)

    def forward(self, *args: Any, **kwargs: Any) -> Any:
        return self.decode(*args, **kwargs)