import torch import torch.nn as nn from models.unet_2d_condition import UNet2DConditionModel from models.unet_3d import UNet3DConditionModel class ChampModel(nn.Module): def __init__( self, reference_unet: UNet2DConditionModel, denoising_unet: UNet3DConditionModel, reference_control_writer, reference_control_reader, guidance_encoder_group, ): super().__init__() self.reference_unet = reference_unet self.denoising_unet = denoising_unet self.reference_control_writer = reference_control_writer self.reference_control_reader = reference_control_reader self.guidance_types = [] self.guidance_input_channels = [] for guidance_type, guidance_module in guidance_encoder_group.items(): setattr(self, f"guidance_encoder_{guidance_type}", guidance_module) self.guidance_types.append(guidance_type) self.guidance_input_channels.append(guidance_module.guidance_input_channels) def forward( self, noisy_latents, timesteps, ref_image_latents, clip_image_embeds, multi_guidance_cond, uncond_fwd: bool = False, ): guidance_cond_group = torch.split( multi_guidance_cond, self.guidance_input_channels, dim=1 ) guidance_fea_lst = [] for guidance_idx, guidance_cond in enumerate(guidance_cond_group): guidance_encoder = getattr( self, f"guidance_encoder_{self.guidance_types[guidance_idx]}" ) guidance_fea = guidance_encoder(guidance_cond) guidance_fea_lst += [guidance_fea] guidance_fea = torch.stack(guidance_fea_lst, dim=0).sum(0) if not uncond_fwd: ref_timesteps = torch.zeros_like(timesteps) self.reference_unet( ref_image_latents, ref_timesteps, encoder_hidden_states=clip_image_embeds, return_dict=False, ) self.reference_control_reader.update(self.reference_control_writer) model_pred = self.denoising_unet( noisy_latents, timesteps, guidance_fea=guidance_fea, encoder_hidden_states=clip_image_embeds, ).sample return model_pred