SubstanceSHIFT/SeedVR2-3B
0
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 