"""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)