Team Ai
Apppublic

modelscope/DiffSynth-Painter

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
14likes
dancer.py234 linesDownload Raw Back to pipelines
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