SolPix / solpix /autoencoder.py
j0no12's picture
Name the SolPix transformer and decoder adapter
c0f018c verified
Raw History Blame Contribute Delete
1.25 kB
"""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)