AINovice2005's picture
Update block.py
44ad3d9 verified
Raw
History Blame Contribute Delete
1.87 kB
from diffusers.modular_pipelines import ModularPipelineBlocks, PipelineState
from diffusers.modular_pipelines.modular_pipeline_utils import (
ComponentSpec,
InputParam,
OutputParam,
)
from .modeling_prunavae import PrunaAutoencoderKLLTX2Video
class LoadPrunaVAE(ModularPipelineBlocks):
model_name = "PrunaVAED"
def __init__(
self,
pretrained_model_name_or_path: str = "AINovice2005/pruna-vaed-modular-diffusers",
subfolder: str | None = None,
revision: str | None = None,
variant: str | None = None,
):
super().__init__()
self.pretrained_model_name_or_path = pretrained_model_name_or_path
self.subfolder = subfolder
self.revision = revision
self.variant = variant
@property
def description(self) -> str:
return (
"Loads a PrunaAutoencoderKLLTX2Video from the Hugging Face Hub "
"and exposes it as `components.vae`."
)
@property
def expected_components(self):
return [
ComponentSpec(
"vae",
PrunaAutoencoderKLLTX2Video,
pretrained_model_name_or_path=self.pretrained_model_name_or_path,
subfolder=self.subfolder,
revision=self.revision,
variant=self.variant,
)
]
@property
def inputs(self):
return []
@property
def intermediate_outputs(self):
return []
def __call__(self, components, state):
# Nothing to compute here -- this block's only role is to make # `components.vae`
#(a PrunaAutoencoderKLLTX2Video) available to # every block downstream in the pipeline.
#Blocks that actually # need it (e.g. a decode step) read it directly off `components`, # not off `state`.
return components, state