CoolFace
Apppublic

SubstanceSHIFT/SeedVR2-3B

sourceHugging Faceapache-2.0updated 1y agoView on Hugging Face
0likes
inference_seedvr2_7b.py322 linesDownload Raw Back to projects
1# // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates2# //3# // Licensed under the Apache License, Version 2.0 (the "License");4# // you may not use this file except in compliance with the License.5# // You may obtain a copy of the License at6# //7# //     http://www.apache.org/licenses/LICENSE-2.08# //9# // Unless required by applicable law or agreed to in writing, software10# // distributed under the License is distributed on an "AS IS" BASIS,11# // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.12# // See the License for the specific language governing permissions and13# // limitations under the License.14 15import os16import torch17import mediapy18from einops import rearrange19from omegaconf import OmegaConf20print(os.getcwd())21import datetime22from tqdm import tqdm23from models.dit import na24import gc25 26from data.image.transforms.divisible_crop import DivisibleCrop27from data.image.transforms.na_resize import NaResize28from data.video.transforms.rearrange import Rearrange29if os.path.exists("./projects/video_diffusion_sr/color_fix.py"):30    from projects.video_diffusion_sr.color_fix import wavelet_reconstruction31    use_colorfix=True32else:33    use_colorfix = False34    print('Note!!!!!! Color fix is not avaliable!')35from torchvision.transforms import Compose, Lambda, Normalize36from torchvision.io.video import read_video37 38 39from common.distributed import (40    get_device,41    init_torch,42)43 44from common.distributed.advanced import (45    get_data_parallel_rank,46    get_data_parallel_world_size,47    get_sequence_parallel_rank,48    get_sequence_parallel_world_size,49    init_sequence_parallel,50)51 52from projects.video_diffusion_sr.infer import VideoDiffusionInfer53from common.config import load_config54from common.distributed.ops import sync_data55from common.seed import set_seed56from common.partition import partition_by_groups, partition_by_size57import argparse58 59def configure_sequence_parallel(sp_size):60    if sp_size > 1:61        init_sequence_parallel(sp_size)62 63def configure_runner(sp_size):64    config_path = os.path.join('./configs_7b', 'main.yaml')65    config = load_config(config_path)66    runner = VideoDiffusionInfer(config)67    OmegaConf.set_readonly(runner.config, False)68    69    init_torch(cudnn_benchmark=False, timeout=datetime.timedelta(seconds=3600))70    configure_sequence_parallel(sp_size)71    runner.configure_dit_model(device="cuda", checkpoint='./ckpts/seedvr2_ema_7b.pth')72    runner.configure_vae_model()73    # Set memory limit.74    if hasattr(runner.vae, "set_memory_limit"):75        runner.vae.set_memory_limit(**runner.config.vae.memory_limit)76    return runner77 78def generation_step(runner, text_embeds_dict, cond_latents):79    def _move_to_cuda(x):80        return [i.to(get_device()) for i in x]81 82    noises = [torch.randn_like(latent) for latent in cond_latents]83    aug_noises = [torch.randn_like(latent) for latent in cond_latents]84    print(f"Generating with noise shape: {noises[0].size()}.")85    noises, aug_noises, cond_latents = sync_data((noises, aug_noises, cond_latents), 0)86    noises, aug_noises, cond_latents = list(87        map(lambda x: _move_to_cuda(x), (noises, aug_noises, cond_latents))88    )89    cond_noise_scale = 0.090 91    def _add_noise(x, aug_noise):92        t = (93            torch.tensor([1000.0], device=get_device())94            * cond_noise_scale95        )96        shape = torch.tensor(x.shape[1:], device=get_device())[None]97        t = runner.timestep_transform(t, shape)98        print(99            f"Timestep shifting from"100            f" {1000.0 * cond_noise_scale} to {t}."101        )102        x = runner.schedule.forward(x, aug_noise, t)103        return x104 105    conditions = [106        runner.get_condition(107            noise,108            task="sr",109            latent_blur=_add_noise(latent_blur, aug_noise),110        )111        for noise, aug_noise, latent_blur in zip(noises, aug_noises, cond_latents)112    ]113 114    with torch.no_grad(), torch.autocast("cuda", torch.bfloat16, enabled=True):115        video_tensors = runner.inference(116            noises=noises,117            conditions=conditions,118            dit_offload=True,119            **text_embeds_dict,120        )121 122    samples = [123        (124            rearrange(video[:, None], "c t h w -> t c h w")125            if video.ndim == 3126            else rearrange(video, "c t h w -> t c h w")127        )128        for video in video_tensors129    ]130    del video_tensors131 132    return samples133 134def generation_loop(runner, video_path='./test_videos', output_dir='./results', batch_size=1, cfg_scale=1.0, cfg_rescale=0.0, sample_steps=1, seed=666, res_h=1280, res_w=720, sp_size=1):135 136    def _build_pos_and_neg_prompt():137        # read positive prompt138        positive_text = "Cinematic, High Contrast, highly detailed, taken using a Canon EOS R camera, \139        hyper detailed photo - realistic maximum detail, 32k, Color Grading, ultra HD, extreme meticulous detailing, \140        skin pore detailing, hyper sharpness, perfect without deformations."141        # read negative prompt142        negative_text = "painting, oil painting, illustration, drawing, art, sketch, oil painting, cartoon, \143        CG Style, 3D render, unreal engine, blurring, dirty, messy, worst quality, low quality, frames, watermark, \144        signature, jpeg artifacts, deformed, lowres, over-smooth"145        return positive_text, negative_text146 147    def _build_test_prompts(video_path):148        positive_text, negative_text = _build_pos_and_neg_prompt()149        original_videos = []150        prompts = {}151        video_list = os.listdir(video_path)152        for f in video_list:153            if f.endswith(".mp4"):154                original_videos.append(f)155                prompts[f] = positive_text156        print(f"Total prompts to be generated: {len(original_videos)}")157        return original_videos, prompts, negative_text158 159    def _extract_text_embeds():160        # Text encoder forward.161        positive_prompts_embeds = []162        for texts_pos in tqdm(original_videos_local):163            text_pos_embeds = torch.load('pos_emb.pt')164            text_neg_embeds = torch.load('neg_emb.pt')165 166            positive_prompts_embeds.append(167                {"texts_pos": [text_pos_embeds], "texts_neg": [text_neg_embeds]}168            )169        gc.collect()170        torch.cuda.empty_cache()171        return positive_prompts_embeds172 173    def cut_videos(videos, sp_size):174        t = videos.size(1)175        if t <= 4 * sp_size:176            print(f"Cut input video size: {videos.size()}")177            padding = [videos[:, -1].unsqueeze(1)] * (4 * sp_size - t + 1)178            padding = torch.cat(padding, dim=1)179            videos = torch.cat([videos, padding], dim=1)180            return videos181        if (t - 1) % (4 * sp_size) == 0:182            return videos183        else:184            padding = [videos[:, -1].unsqueeze(1)] * (185                4 * sp_size - ((t - 1) % (4 * sp_size))186            )187            padding = torch.cat(padding, dim=1)188            videos = torch.cat([videos, padding], dim=1)189            assert (videos.size(1) - 1) % (4 * sp_size) == 0190            return videos191 192    # classifier-free guidance193    runner.config.diffusion.cfg.scale = cfg_scale194    runner.config.diffusion.cfg.rescale = cfg_rescale195    # sampling steps196    runner.config.diffusion.timesteps.sampling.steps = sample_steps197    runner.configure_diffusion()198 199    # set random seed200    set_seed(seed, same_across_ranks=True)201    os.makedirs(output_dir, exist_ok=True)202    tgt_path = output_dir203 204    # get test prompts205    original_videos, _, _ = _build_test_prompts(video_path)206 207    # divide the prompts into different groups208    original_videos_group = partition_by_groups(209        original_videos,210        get_data_parallel_world_size() // get_sequence_parallel_world_size(),211    )212    # store prompt mapping213    original_videos_local = original_videos_group[214        get_data_parallel_rank() // get_sequence_parallel_world_size()215    ]216    original_videos_local = partition_by_size(original_videos_local, batch_size)217 218    # pre-extract the text embeddings219    positive_prompts_embeds = _extract_text_embeds()220 221    video_transform = Compose(222        [223            NaResize(224                resolution=(225                    res_h * res_w226                )227                ** 0.5,228                mode="area",229                # Upsample image, model only trained for high res.230                downsample_only=False,231            ),232            Lambda(lambda x: torch.clamp(x, 0.0, 1.0)),233            DivisibleCrop((16, 16)),234            Normalize(0.5, 0.5),235            Rearrange("t c h w -> c t h w"),236        ]237    )238 239    # generation loop240    for videos, text_embeds in tqdm(zip(original_videos_local, positive_prompts_embeds)):241        # read condition latents242        cond_latents = []243        for video in videos:244            video = (245                read_video(246                   os.path.join(video_path, video), output_format="TCHW"247                )[0]248                / 255.0249            )250            print(f"Read video size: {video.size()}")251            cond_latents.append(video_transform(video.to(get_device())))252 253        ori_lengths = [video.size(1) for video in cond_latents]254        input_videos = cond_latents255        cond_latents = [cut_videos(video, sp_size) for video in cond_latents]256 257        runner.dit.to("cpu")258        print(f"Encoding videos: {list(map(lambda x: x.size(), cond_latents))}")259        runner.vae.to(get_device())260        cond_latents = runner.vae_encode(cond_latents)261        runner.vae.to("cpu")262        runner.dit.to(get_device())263 264        for i, emb in enumerate(text_embeds["texts_pos"]):265            text_embeds["texts_pos"][i] = emb.to(get_device())266        for i, emb in enumerate(text_embeds["texts_neg"]):267            text_embeds["texts_neg"][i] = emb.to(get_device())268 269        samples = generation_step(runner, text_embeds, cond_latents=cond_latents)270        runner.dit.to("cpu")271        del cond_latents272 273        # dump samples to the output directory274        if get_sequence_parallel_rank() == 0:275            for path, input, sample, ori_length in zip(276                videos, input_videos, samples, ori_lengths277            ):278                if ori_length < sample.shape[0]:279                    sample = sample[:ori_length]280                filename = os.path.join(tgt_path, os.path.basename(path))281                # color fix282                input = (283                    rearrange(input[:, None], "c t h w -> t c h w")284                    if input.ndim == 3285                    else rearrange(input, "c t h w -> t c h w")286                )287                if use_colorfix:288                    sample = wavelet_reconstruction(289                        sample.to("cpu"), input[: sample.size(0)].to("cpu")290                    )291                else:292                    sample = sample.to("cpu")293                sample = (294                    rearrange(sample[:, None], "t c h w -> t h w c")295                    if sample.ndim == 3296                    else rearrange(sample, "t c h w -> t h w c")297                )298                sample = sample.clip(-1, 1).mul_(0.5).add_(0.5).mul_(255).round()299                sample = sample.to(torch.uint8).numpy()300 301                if sample.shape[0] == 1:302                    mediapy.write_image(filename, sample.squeeze(0))303                else:304                    mediapy.write_video(305                        filename, sample, fps=24306                    )307        gc.collect()308        torch.cuda.empty_cache()309 310if __name__ == "__main__":311    parser = argparse.ArgumentParser() 312    parser.add_argument("--video_path", type=str, default="./test_videos")313    parser.add_argument("--output_dir", type=str, default="./results")314    parser.add_argument("--seed", type=int, default=666)315    parser.add_argument("--res_h", type=int, default=720)316    parser.add_argument("--res_w", type=int, default=1280)317    parser.add_argument("--sp_size", type=int, default=1)318    args = parser.parse_args()319 320    runner = configure_runner(args.sp_size)321    generation_loop(runner, **vars(args))322