CoolFace
Apppublic

Kleinhe/SemanticBoost

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
double_take.py270 linesDownload Raw Back to motion
1from copy import deepcopy2import torch3import pandas as pd4import numpy as np5 6def pad_sample_with_zeros(sample, max_len=250):7    # pad inp, change lenghts, and pad is transition8    seq_len, n_feats = sample.shape9    len_to_pad = max_len - seq_len10    np.zeros_like(sample)11    sample_padding = np.zeros((len_to_pad, n_feats))12    sample = np.concatenate((sample, sample_padding))13    return sample14 15def split2subs(motions, step_sizes, batch_size, blend_len, max_motion_length):16    #### motions [1, 263, 1, Nlength] -> [263, Nlength] -> [NLength, 263]17    new_motions = []18    new_lengths = []19    new_motions.append(pad_sample_with_zeros(motions[..., :step_sizes[0] - blend_len].squeeze().permute(1, 0).cpu().numpy(), max_motion_length))20    new_lengths.append(step_sizes[0] - blend_len)21    for i in range(1, batch_size-1):22        curr = pad_sample_with_zeros(motions[..., step_sizes[i-1]-blend_len:step_sizes[i]-blend_len].squeeze().permute(1, 0).cpu().numpy(), max_motion_length)23        new_motions.append(curr)24        new_lengths.append(step_sizes[i] - step_sizes[i-1])25 26    new_motions.append(pad_sample_with_zeros(motions[..., step_sizes[-1]-blend_len:].squeeze().permute(1, 0).cpu().numpy(), max_motion_length))27    new_lengths.append(step_sizes[-1]-step_sizes[-2]+blend_len)28 29    new_motions = np.stack(new_motions, axis=0)30    new_motions = torch.from_numpy(new_motions)31    new_lengths = np.stack(new_lengths, axis=0)32    new_lengths = torch.from_numpy(new_lengths).long()33    return new_motions, new_lengths34 35 36def unfold_sample_arb_len(sample, handshake_size, step_sizes, final_n_frames, model_kwargs):37    old_sample = deepcopy(sample)38    new_shape = list(old_sample.shape)39    new_shape[0] = 140    new_shape[-1] = final_n_frames41    sample = torch.zeros(new_shape, dtype=sample.dtype, device=sample.device)42    sample[0, :, :, :model_kwargs['y']['lengths'][0]] = old_sample[0, :, :, :model_kwargs['y']['lengths'][0]]43    for sample_i, len_i in enumerate(step_sizes):44        if sample_i == 0:45            continue46        start = step_sizes[sample_i-1]47        sample[0, :, :, start:len_i] = old_sample[sample_i, :, :, handshake_size:model_kwargs['y']['lengths'][sample_i]]48    return sample49 50 51def double_take_arb_len(diffusion, model, model_kwargs, n_frames, blend_len=10, handshake_size=20, device="cpu", progress=True):52    sample_fn = diffusion.p_sample_loop53    blend_len = blend_len54    handshake_size = handshake_size55 56    batch_size = len(model_kwargs['y']['text'])57 58    # Unfolding - orig59    sample = sample_fn(60        model,61        (batch_size, model.njoints, model.nfeats, n_frames),62        clip_denoised=False,63        model_kwargs=model_kwargs,64        skip_timesteps=0,  # 0 is the default value - i.e. don't skip any step65        init_image=None,66        progress=progress,67        dump_steps=None,68        noise=None,69        const_noise=False,70        unfolding_handshake=handshake_size,71    )72 73    model_kwargs['y']['scale'] = torch.ones(batch_size-1, device=device) * 074    sample = sample["output"]       #### [5, 263, 1 196]75 76    '''77    1. 替换 sample 78    2. model_kwargs['y']['lengths']79    '''80 81    new_sample_seq_len = (sample.shape[-1] - 2 * handshake_size) * 2 + handshake_size82 83    bs, feats, joints, seq_len = sample.shape84    new_sample = torch.zeros((bs-1, feats, joints, new_sample_seq_len), dtype=sample.dtype, device=sample.device)85 86    generated_motion = []87    right_constraint = []88    left_constraint = []89 90    for ii in range(bs):        #####  按左中右拆分 Motion91        generated_motion.append(deepcopy(sample[ii, :, :, handshake_size: model_kwargs['y']['lengths'][ii]-handshake_size])) # w/o start and end92        left_constraint.append(deepcopy(sample[ii, :, :, :handshake_size]))  # left side93        right_constraint.append(deepcopy(sample[ii, :, :, model_kwargs['y']['lengths'][ii] - handshake_size: model_kwargs['y']['lengths'][ii]]))94 95    buffer = []     #### 存放剩下的动作部分的长度,也就是 generated_motion 的长度96    for ii in range(bs):97        buffer.append(int(model_kwargs['y']['lengths'][ii]) - 2*handshake_size)98    for ii in range(bs - 1):  # run over bs, 把 N句话 合并成 N-1 句话,新 motion 的组成 [gm[i-1], right[i-1], gm[i]], 长度是 2 * gm_length + hand_size99        new_sample[ii, :, :, :buffer[ii]] = generated_motion[ii]100        new_sample[ii, :, :, buffer[ii]: buffer[ii]+handshake_size] = right_constraint[ii] # add transition101        new_sample[ii, :, :, buffer[ii]+handshake_size : buffer[ii]+handshake_size+buffer[ii+1]] = generated_motion[ii + 1]102 103    # "in between"104    model_kwargs['y']['inpainted_motion'] = new_sample105    model_kwargs['y']['inpainting_mask'] = torch.ones_like(new_sample, dtype=torch.float,106                                                            device=new_sample.device)107 108    for ii in range(bs - 1):  # run over bs109        if blend_len >= 2:110            '''111            渐变混合112            1. 在左边 gm[i-1] 靠后 blend_len 的区域,渐变地保留原本的内容113            2. 在右边 gm[i] 靠前的 blend_len 的区域,渐变保留原本的内容114            3. 似乎是 right 的部分完全保留,也就是用前一个动作的结束座位后一个动作的开头 115            '''116 117            model_kwargs['y']['inpainting_mask'][ii, :, :, buffer[ii] - blend_len: buffer[ii]] = \118                torch.arange(0.85, 0.0, -0.85 / int(blend_len))119            model_kwargs['y']['inpainting_mask'][ii, :, :, buffer[ii] + handshake_size: buffer[ii] + handshake_size + blend_len] = \120                torch.arange(0.0, 0.85, 0.85 / int(blend_len))121 122    model_kwargs['y']['uncond'] = 1.0       ### 混合多段语意后,cond 没什么意义,而且需要生成的内容很少123    model_kwargs['y']['text'] = model_kwargs['y']['text'][:bs-1]124    sample_fn = diffusion.p_sample_loop  # double take sample function125    n_frames = new_sample_seq_len126    orig_lens = deepcopy(model_kwargs['y']['lengths'])127    for ii in range (len(model_kwargs['y']['lengths'])-1):128        model_kwargs['y']['lengths'][ii] = model_kwargs['y']['lengths'][ii] + model_kwargs['y']['lengths'][ii+1] - 3*handshake_size129    model_kwargs['y']['lengths'] = model_kwargs['y']['lengths'][:-1]130 131    double_take_sample = sample_fn(132        model,133        (batch_size-1, model.njoints, model.nfeats, n_frames),134        clip_denoised=False,135        model_kwargs=model_kwargs,136        skip_timesteps=0,  # 0 is the default value - i.e. don't skip any step137        init_image=new_sample, #TODO!! check if plausible or not!138        progress=progress,139        dump_steps=None,140        noise=None,141        const_noise=False,142    )143    double_take_sample = double_take_sample["output"]144    model_kwargs['y']['lengths'] = orig_lens145    # rebuild_orig:146    rebuild_sample = torch.zeros_like(sample)147 148    '''149    sample -> left + motion + right150    double_take_sample -> motion1 + blend + hand + blend + motion2, 其中长度表示 : motion1 + blend = motion2 + blend = motion151    '''152 153    transitions, right_side, left_side = [], [], []154    for ii in range(bs - 1):  # run over bs155        transitions.append(double_take_sample[ii, :, :, buffer[ii]: buffer[ii]+handshake_size])156        right_side.append(double_take_sample[ii, :, :, buffer[ii] + handshake_size: buffer[ii] + handshake_size + blend_len]) # M1 blending..157        left_side.append(double_take_sample[ii, :, :, buffer[ii] - blend_len:buffer[ii]]) # M0 blending...158 159        '''160        translation 储存的是 hand161        right_side 存右边的 blend162        left_side 村左边的 blend163        '''164 165 166    rebuild_sample[0, :, :, :handshake_size] = left_constraint[0] # Fill missing167    rebuild_sample[-1, :, :, buffer[-1]+handshake_size: buffer[-1]+2*handshake_size] = right_constraint[-1] # Fill missing168 169    '''170    展开 double take 的结果, 还原会原本的状态,即 left + motion + right171    '''172 173    for ii in range(bs - 1):174        rebuild_sample[ii + 1, :, :, :handshake_size] = transitions[ii]175        rebuild_sample[ii, :, :, handshake_size: buffer[ii]+handshake_size] = generated_motion[ii]176        rebuild_sample[ii, :, :, buffer[ii]+handshake_size: buffer[ii]+2*handshake_size] = transitions[ii]      #### motion1 的 right = motion2 的 left177        rebuild_sample[ii, :, :, handshake_size + buffer[ii]-blend_len: handshake_size + buffer[ii]] = left_side[ii]178        # if ii > 0:179    rebuild_sample[-1, :, :, handshake_size: buffer[-1] + handshake_size] = generated_motion[-1]180    for ii in range(bs - 1):181        rebuild_sample[ii+1, :, :, handshake_size:handshake_size + blend_len] = right_side[ii]182 183    double_take_sample = deepcopy(rebuild_sample)184 185    return double_take_sample186 187def double_take(prompt=None, path=None, num_repetitions=1, model=None, diffusion=None, handshake_size=20, blend_len=10, default_length=196, guidance_param=2.5, device="cpu", progress=True):188    assert model is not None189    assert diffusion is not None190    if prompt is not None:191        texts = prompt.split("|")192        num_samples = len(texts)193        length = []194        captions = []195        for i in range(len(texts)):196            nframes = texts[i].split(",")[0]197            try:198                nframes = int(nframes)199                curr_text = texts[i].split(",")[1::]200                curr_text = ",".join(curr_text)201            except:202                nframes = default_length203                curr_text = texts[i]204 205            captions.append(curr_text)206            length.append(nframes)207 208        model_kwargs = {'y': {209            'mask': torch.ones((len(texts), 1, 1, default_length)), # 196 is humanml max frames number210            'lengths': torch.tensor(length),211            'text': captions,212            'tokens': [''],213            'scale': torch.ones(len(texts))*guidance_param214        }}215    elif path.split(".")[-1] == "csv":216        df = pd.read_csv(path)217        num_samples = len(list(df['text']))  218        model_kwargs = {'y': {219            'mask': torch.ones((len(list(df['text'])), 1, 1, default_length)), #196 is humanml max frames number220            'lengths': torch.tensor(list(df['length'])),221            'text': list(df['text']),222            'tokens': [''],223            'scale': torch.ones(len(list(df['text'])))*guidance_param224        }}  225    elif path.split(".")[-1] == "txt":226        with open(path, 'r') as fr:227            texts = fr.readlines()228        texts = [s.replace('\n', '') for s in texts]229        num_samples = len(texts)      230        model_kwargs = {'y': {231            'mask': torch.ones((len(texts), 1, 1, default_length)), # 196 is humanml max frames number232            'lengths': torch.tensor([default_length]*len(texts)),233            'text': texts,234            'tokens': [''],235            'scale': torch.ones(len(texts))*guidance_param236        }}237 238    all_motions = []239 240    for rep_i in range(num_repetitions):241        if guidance_param != 1:242            model_kwargs['y']['scale'] = torch.ones(num_samples, device=device) * guidance_param243        model_kwargs['y'] = {key: val.to(device) if torch.is_tensor(val) else val for key, val in model_kwargs['y'].items()}244 245        max_arb_len = model_kwargs['y']['lengths'].max()246        min_arb_len = 2 * handshake_size + 2*blend_len + 10247 248        for ii, len_s in enumerate(model_kwargs['y']['lengths']):249            if len_s > max_arb_len:250                model_kwargs['y']['lengths'][ii] = max_arb_len251            if len_s < min_arb_len:252                model_kwargs['y']['lengths'][ii] = min_arb_len253 254        sample = double_take_arb_len(diffusion, model, model_kwargs, max_arb_len, blend_len, handshake_size, device, progress=progress)    255        step_sizes = np.zeros(len(model_kwargs['y']['lengths']), dtype=int)256        for ii, len_i in enumerate(model_kwargs['y']['lengths']):257            if ii == 0:258                step_sizes[ii] = len_i259                continue260            step_sizes[ii] = step_sizes[ii-1] + len_i - handshake_size261 262        final_n_frames = step_sizes[-1]263        sample = unfold_sample_arb_len(sample, handshake_size, step_sizes, final_n_frames, model_kwargs)264 265        all_motions.append(sample)266    267    all_motions = torch.cat(all_motions, dim=0)268    return all_motions, step_sizes269 270