| 13 | logger = logging.getLogger(__name__) |
| 14 | |
| 15 | class MimicMotionModel(torch.nn.Module): |
| 16 | def __init__(self, base_model_path): |
| 17 | """construnct base model components and load pretrained svd model except pose-net |
| 18 | Args: |
| 19 | base_model_path (str): pretrained svd model path |
| 20 | """ |
| 21 | super().__init__() |
| 22 | self.unet = UNetSpatioTemporalConditionModel.from_config( |
| 23 | UNetSpatioTemporalConditionModel.load_config(base_model_path, subfolder="unet")) |
| 24 | self.vae = AutoencoderKLTemporalDecoder.from_pretrained( |
| 25 | base_model_path, subfolder="vae", torch_dtype=torch.float16, variant="fp16") |
| 26 | self.image_encoder = CLIPVisionModelWithProjection.from_pretrained( |
| 27 | base_model_path, subfolder="image_encoder", torch_dtype=torch.float16, variant="fp16") |
| 28 | self.noise_scheduler = EulerDiscreteScheduler.from_pretrained( |
| 29 | base_model_path, subfolder="scheduler") |
| 30 | self.feature_extractor = CLIPImageProcessor.from_pretrained( |
| 31 | base_model_path, subfolder="feature_extractor") |
| 32 | # pose_net |
| 33 | self.pose_net = PoseNet(noise_latent_channels=self.unet.config.block_out_channels[0]) |
| 34 | |
| 35 | def create_pipeline(infer_config, device): |
| 36 | """create mimicmotion pipeline and load pretrained weight |