modelscope/DiffSynth-Painter
14
1import torch2from ..models import SDUNet, SDMotionModel, SDXLUNet, SDXLMotionModel3from ..models.sd_unet import PushBlock, PopBlock4from ..controlnets import MultiControlNetManager5 6 7def lets_dance(8 unet: SDUNet,9 motion_modules: SDMotionModel = None,10 controlnet: MultiControlNetManager = None,11 sample = None,12 timestep = None,13 encoder_hidden_states = None,14 ipadapter_kwargs_list = {},15 controlnet_frames = None,16 unet_batch_size = 1,17 controlnet_batch_size = 1,18 cross_frame_attention = False,19 tiled=False,20 tile_size=64,21 tile_stride=32,22 device = "cuda",23 vram_limit_level = 0,24):25 # 0. Text embedding alignment (only for video processing)26 if encoder_hidden_states.shape[0] != sample.shape[0]:27 encoder_hidden_states = encoder_hidden_states.repeat(sample.shape[0], 1, 1, 1)28 29 # 1. ControlNet30 # This part will be repeated on overlapping frames if animatediff_batch_size > animatediff_stride.31 # I leave it here because I intend to do something interesting on the ControlNets.32 controlnet_insert_block_id = 3033 if controlnet is not None and controlnet_frames is not None:34 res_stacks = []35 # process controlnet frames with batch36 for batch_id in range(0, sample.shape[0], controlnet_batch_size):37 batch_id_ = min(batch_id + controlnet_batch_size, sample.shape[0])38 res_stack = controlnet(39 sample[batch_id: batch_id_],40 timestep,41 encoder_hidden_states[batch_id: batch_id_],42 controlnet_frames[:, batch_id: batch_id_],43 tiled=tiled, tile_size=tile_size, tile_stride=tile_stride44 )45 if vram_limit_level >= 1:46 res_stack = [res.cpu() for res in res_stack]47 res_stacks.append(res_stack)48 # concat the residual49 additional_res_stack = []50 for i in range(len(res_stacks[0])):51 res = torch.concat([res_stack[i] for res_stack in res_stacks], dim=0)52 additional_res_stack.append(res)53 else:54 additional_res_stack = None55 56 # 2. time57 time_emb = unet.time_proj(timestep).to(sample.dtype)58 time_emb = unet.time_embedding(time_emb)59 60 # 3. pre-process61 height, width = sample.shape[2], sample.shape[3]62 hidden_states = unet.conv_in(sample)63 text_emb = encoder_hidden_states64 res_stack = [hidden_states.cpu() if vram_limit_level>=1 else hidden_states]65 66 # 4. blocks67 for block_id, block in enumerate(unet.blocks):68 # 4.1 UNet69 if isinstance(block, PushBlock):70 hidden_states, time_emb, text_emb, res_stack = block(hidden_states, time_emb, text_emb, res_stack)71 if vram_limit_level>=1:72 res_stack[-1] = res_stack[-1].cpu()73 elif isinstance(block, PopBlock):74 if vram_limit_level>=1:75 res_stack[-1] = res_stack[-1].to(device)76 hidden_states, time_emb, text_emb, res_stack = block(hidden_states, time_emb, text_emb, res_stack)77 else:78 hidden_states_input = hidden_states79 hidden_states_output = []80 for batch_id in range(0, sample.shape[0], unet_batch_size):81 batch_id_ = min(batch_id + unet_batch_size, sample.shape[0])82 hidden_states, _, _, _ = block(83 hidden_states_input[batch_id: batch_id_],84 time_emb,85 text_emb[batch_id: batch_id_],86 res_stack,87 cross_frame_attention=cross_frame_attention,88 ipadapter_kwargs_list=ipadapter_kwargs_list.get(block_id, {}),89 tiled=tiled, tile_size=tile_size, tile_stride=tile_stride90 )91 hidden_states_output.append(hidden_states)92 hidden_states = torch.concat(hidden_states_output, dim=0)93 # 4.2 AnimateDiff94 if motion_modules is not None:95 if block_id in motion_modules.call_block_id:96 motion_module_id = motion_modules.call_block_id[block_id]97 hidden_states, time_emb, text_emb, res_stack = motion_modules.motion_modules[motion_module_id](98 hidden_states, time_emb, text_emb, res_stack,99 batch_size=1100 )101 # 4.3 ControlNet102 if block_id == controlnet_insert_block_id and additional_res_stack is not None:103 hidden_states += additional_res_stack.pop().to(device)104 if vram_limit_level>=1:105 res_stack = [(res.to(device) + additional_res.to(device)).cpu() for res, additional_res in zip(res_stack, additional_res_stack)]106 else:107 res_stack = [res + additional_res for res, additional_res in zip(res_stack, additional_res_stack)]108 109 # 5. output110 hidden_states = unet.conv_norm_out(hidden_states)111 hidden_states = unet.conv_act(hidden_states)112 hidden_states = unet.conv_out(hidden_states)113 114 return hidden_states115 116 117 118 119def lets_dance_xl(120 unet: SDXLUNet,121 motion_modules: SDXLMotionModel = None,122 controlnet: MultiControlNetManager = None,123 sample = None,124 add_time_id = None,125 add_text_embeds = None,126 timestep = None,127 encoder_hidden_states = None,128 ipadapter_kwargs_list = {},129 controlnet_frames = None,130 unet_batch_size = 1,131 controlnet_batch_size = 1,132 cross_frame_attention = False,133 tiled=False,134 tile_size=64,135 tile_stride=32,136 device = "cuda",137 vram_limit_level = 0,138):139 # 0. Text embedding alignment (only for video processing)140 if encoder_hidden_states.shape[0] != sample.shape[0]:141 encoder_hidden_states = encoder_hidden_states.repeat(sample.shape[0], 1, 1, 1)142 143 # 1. ControlNet144 controlnet_insert_block_id = 22145 if controlnet is not None and controlnet_frames is not None:146 res_stacks = []147 # process controlnet frames with batch148 for batch_id in range(0, sample.shape[0], controlnet_batch_size):149 batch_id_ = min(batch_id + controlnet_batch_size, sample.shape[0])150 res_stack = controlnet(151 sample[batch_id: batch_id_],152 timestep,153 encoder_hidden_states[batch_id: batch_id_],154 controlnet_frames[:, batch_id: batch_id_],155 add_time_id=add_time_id,156 add_text_embeds=add_text_embeds,157 tiled=tiled, tile_size=tile_size, tile_stride=tile_stride,158 unet=unet, # for Kolors, some modules in ControlNets will be replaced.159 )160 if vram_limit_level >= 1:161 res_stack = [res.cpu() for res in res_stack]162 res_stacks.append(res_stack)163 # concat the residual164 additional_res_stack = []165 for i in range(len(res_stacks[0])):166 res = torch.concat([res_stack[i] for res_stack in res_stacks], dim=0)167 additional_res_stack.append(res)168 else:169 additional_res_stack = None170 171 # 2. time172 t_emb = unet.time_proj(timestep).to(sample.dtype)173 t_emb = unet.time_embedding(t_emb)174 175 time_embeds = unet.add_time_proj(add_time_id)176 time_embeds = time_embeds.reshape((add_text_embeds.shape[0], -1))177 add_embeds = torch.concat([add_text_embeds, time_embeds], dim=-1)178 add_embeds = add_embeds.to(sample.dtype)179 add_embeds = unet.add_time_embedding(add_embeds)180 181 time_emb = t_emb + add_embeds182 183 # 3. pre-process184 height, width = sample.shape[2], sample.shape[3]185 hidden_states = unet.conv_in(sample)186 text_emb = encoder_hidden_states if unet.text_intermediate_proj is None else unet.text_intermediate_proj(encoder_hidden_states)187 res_stack = [hidden_states]188 189 # 4. blocks190 for block_id, block in enumerate(unet.blocks):191 # 4.1 UNet192 if isinstance(block, PushBlock):193 hidden_states, time_emb, text_emb, res_stack = block(hidden_states, time_emb, text_emb, res_stack)194 if vram_limit_level>=1:195 res_stack[-1] = res_stack[-1].cpu()196 elif isinstance(block, PopBlock):197 if vram_limit_level>=1:198 res_stack[-1] = res_stack[-1].to(device)199 hidden_states, time_emb, text_emb, res_stack = block(hidden_states, time_emb, text_emb, res_stack)200 else:201 hidden_states_input = hidden_states202 hidden_states_output = []203 for batch_id in range(0, sample.shape[0], unet_batch_size):204 batch_id_ = min(batch_id + unet_batch_size, sample.shape[0])205 hidden_states, _, _, _ = block(206 hidden_states_input[batch_id: batch_id_],207 time_emb,208 text_emb[batch_id: batch_id_],209 res_stack,210 cross_frame_attention=cross_frame_attention,211 ipadapter_kwargs_list=ipadapter_kwargs_list.get(block_id, {}),212 tiled=tiled, tile_size=tile_size, tile_stride=tile_stride,213 )214 hidden_states_output.append(hidden_states)215 hidden_states = torch.concat(hidden_states_output, dim=0)216 # 4.2 AnimateDiff217 if motion_modules is not None:218 if block_id in motion_modules.call_block_id:219 motion_module_id = motion_modules.call_block_id[block_id]220 hidden_states, time_emb, text_emb, res_stack = motion_modules.motion_modules[motion_module_id](221 hidden_states, time_emb, text_emb, res_stack,222 batch_size=1223 )224 # 4.3 ControlNet225 if block_id == controlnet_insert_block_id and additional_res_stack is not None:226 hidden_states += additional_res_stack.pop().to(device)227 res_stack = [res + additional_res for res, additional_res in zip(res_stack, additional_res_stack)]228 229 # 5. output230 hidden_states = unet.conv_norm_out(hidden_states)231 hidden_states = unet.conv_act(hidden_states)232 hidden_states = unet.conv_out(hidden_states)233 234 return hidden_states