diff --git a/diffuse.py b/diffuse.py index 1a0e337..4e1a7c7 100644 --- a/diffuse.py +++ b/diffuse.py @@ -43,30 +43,18 @@ from .do_run import do_run # !! }} #@title Do the Run! #@markdown `n_batches` ignored with animation modes. -def diffuse(clip_model, clip_vision, args: DiscoDiffusionSettings, batchNum): +def diffuse(model, diffusion, clip_model, clip_vision, args: DiscoDiffusionSettings, batchNum): args.display_rate = 20 #@param{type: 'number'} - args.n_batches = 50 #@param{type: 'number'} if args.animation_mode == 'Video Input': args.steps = args.video_init_steps - #Update Model Settings - timestep_respacing = f'ddim{args.steps}' - diffusion_steps = (1000//args.steps)*args.steps if args.steps < 1000 else args.steps - args.MS.model_config.update({ - 'timestep_respacing': timestep_respacing, - 'diffusion_steps': diffusion_steps, - }) - def move_files(start_num, end_num, old_folder, new_folder): for i in range(start_num, end_num): old_file = old_folder + f'/{args.batch_name}({batchNum})_{i:04}.png' new_file = new_folder + f'/{args.batch_name}({batchNum})_{i:04}.png' os.rename(old_file, new_file) - #@markdown --- - - args.resume_run = False #@param{type: 'boolean'} run_to_resume = 'latest' #@param{type: 'string'} resume_from_frame = 'latest' #@param{type: 'string'} @@ -126,7 +114,7 @@ def diffuse(clip_model, clip_vision, args: DiscoDiffusionSettings, batchNum): seed = int(args.set_seed) args.n_batches = args.n_batches if args.animation_mode == 'None' else 1 - args.max_frames = args.max_frames if args.animation_mode == 'None' else 1 + args.max_frames = args.max_frames if args.animation_mode != 'None' else 1 args.start_frame = start_frame args.seed = seed args.prompts_series = disco_utils.split_prompts(args.text_prompts, args.max_frames) if args.text_prompts else None, @@ -251,28 +239,17 @@ def diffuse(clip_model, clip_vision, args: DiscoDiffusionSettings, batchNum): # args = SimpleNamespace(**args) - device = comfy.model_management.get_torch_device() - - print('Prepping model...') - model, diffusion = create_model_and_diffusion(**args.MS.model_config) - if args.MS.diffusion_model == 'custom': - model.load_state_dict(torch.load(args.MS.custom_path, map_location='cpu')) - else: - model.load_state_dict(torch.load(f'{args.MS.model_path}/{args.MS.get_model_filename(args.MS.diffusion_model)}', map_location='cpu')) - model.requires_grad_(False).eval().to(device) - for name, param in model.named_parameters(): - if 'qkv' in name or 'norm' in name or 'proj' in name: - param.requires_grad_() - if args.MS.model_config['use_fp16']: - model.convert_to_fp16() + results = [] gc.collect() torch.cuda.empty_cache() try: - do_run(diffusion, model, clip_model, clip_vision, args, batchNum) + results = do_run(diffusion, model, clip_model, clip_vision, args, batchNum) except KeyboardInterrupt: pass finally: print('Seed used:', seed) gc.collect() torch.cuda.empty_cache() + + return results diff --git a/do_run.py b/do_run.py index b83f845..8d4fcd5 100644 --- a/do_run.py +++ b/do_run.py @@ -26,6 +26,8 @@ from .CLIP import clip import comfy.model_management import comfy.utils +from comfy.clip_vision import ClipVisionModel +import comfy.sd from . import disco_utils from .make_cutouts import MakeCutouts, MakeCutoutsDango @@ -37,14 +39,22 @@ stop_on_next_loop = False TRANSLATION_SCALE = 1.0/200.0 def encode_text(clip_model, prompt): - _, pooled = clip_model.encode_from_tokens(clip_model.tokenize(prompt), return_pooled=True) - return pooled.float() + if isinstance(clip_model, comfy.sd.CLIP): + # ComfyUI + _, cond_pooled = clip_model.encode_from_tokens(clip_model.tokenize(prompt), return_pooled=True) + return cond_pooled.float() + else: + # OpenAI/OpenClip + device = comfy.model_management.get_torch_device() + return clip_model.encode_text(clip.tokenize([prompt]).to(device)).float() def encode_images(clip_vision, images): - imgs = torch.clip((255. * images), 0, 255).round().int() - inputs = clip_vision.processor(images=imgs, return_tensors="pt") - outputs = clip_vision.model(**inputs) - return outputs.image_embeds.float() + if isinstance(clip_vision, ClipVisionModel): + # ComfyUI + return clip_vision.model(pixel_values=images).image_embeds.float() + else: + # OpenAI/OpenClip + return clip_vision.encode_image(images).float() def do_3d_step(args: DiscoDiffusionSettings, img_filepath, frame_num, midas_model, midas_transform): if args.key_frames: @@ -102,492 +112,521 @@ def id(x): return x +def get_input_resolution(clip_model): + # when using SLIP Base model the dimensions need to be hard coded to avoid AttributeError: 'VisionTransformer' object has no attribute 'input_resolution' + try: + if isinstance(clip_model, ClipVisionModel): + # ComfyUI (Transformers) + return clip_model.model.config.image_size + else: + # OpenAI/OpenClip + return clip_model.visual.input_resolution + except Exception as err: + print("Couldn't find clip vision image size! " + str(err) + " " + str(type(clip_model))) + return 224 + + def do_run(diffusion, model, clip_model, clip_vision, args: DiscoDiffusionSettings, batchNum): - seed = args.seed + global stop_on_next_loop + print(range(args.start_frame, args.max_frames)) pbar = comfy.utils.ProgressBar(diffusion.num_timesteps - args.skip_steps) + midas_model = None + midas_transform = None + midas_net_w = None + midas_net_h = None + midas_resize_mode = None + midas_normalization = None + if (args.animation_mode == "3D") and (args.midas_weight > 0.0): midas_model, midas_transform, midas_net_w, midas_net_h, midas_resize_mode, midas_normalization = init_midas_depth_model( args.midas_depth_model) + results = [] + for frame_num in range(args.start_frame, args.max_frames): if stop_on_next_loop: break - # display.clear_output(wait=True) + results += run_one_frame(diffusion, model, clip_model, clip_vision, args, batchNum, frame_num, pbar, + midas_model, midas_transform, midas_net_w, midas_net_h, midas_resize_mode, midas_normalization) - # Print Frame progress if animation mode is on - # if args.animation_mode != "None": - # batchBar = tqdm(range(args.max_frames), desc="Frames") - # batchBar.n = frame_num - # batchBar.refresh() + return results - # Inits if not video frames - if args.animation_mode != "Video Input": - if args.init_image in ['', 'none', 'None', 'NONE']: - init_image = None + +def run_one_frame(diffusion, model, clip_model, clip_vision, args, batchNum, frame_num, pbar, + midas_model, midas_transform, midas_net_w, midas_net_h, midas_resize_mode, midas_normalization): + global stop_on_next_loop + + # display.clear_output(wait=True) + + # Print Frame progress if animation mode is on + # if args.animation_mode != "None": + # batchBar = tqdm(range(args.max_frames), desc="Frames") + # batchBar.n = frame_num + # batchBar.refresh() + + # Inits if not video frames + if args.animation_mode != "Video Input": + if args.init_image in ['', 'none', 'None', 'NONE']: + init_image = None + else: + init_image = args.init_image + init_scale = args.init_scale + skip_steps = args.skip_steps + + if args.animation_mode == "2D": + if args.key_frames: + angle = args.angle_series[frame_num] + zoom = args.zoom_series[frame_num] + translation_x = args.translation_x_series[frame_num] + translation_y = args.translation_y_series[frame_num] + print( + f'angle: {angle}', + f'zoom: {zoom}', + f'translation_x: {translation_x}', + f'translation_y: {translation_y}', + ) + + if frame_num > 0: + args.seed += 1 + if args.resume_run and frame_num == args.start_frame: + img_0 = cv2.imread( + args.batchFolder+f"/{args.batch_name}({batchNum})_{args.start_frame-1:04}.png") else: - init_image = args.init_image - init_scale = args.init_scale - skip_steps = args.skip_steps + img_0 = cv2.imread('prevFrame.png') + center = (1*img_0.shape[1]//2, 1*img_0.shape[0]//2) + trans_mat = np.float32( + [[1, 0, translation_x], + [0, 1, translation_y]] + ) + rot_mat = cv2.getRotationMatrix2D(center, angle, zoom) + trans_mat = np.vstack([trans_mat, [0, 0, 1]]) + rot_mat = np.vstack([rot_mat, [0, 0, 1]]) + transformation_matrix = np.matmul(rot_mat, trans_mat) + img_0 = cv2.warpPerspective( + img_0, + transformation_matrix, + (img_0.shape[1], img_0.shape[0]), + borderMode=cv2.BORDER_WRAP + ) - if args.animation_mode == "2D": - if args.key_frames: - angle = args.angle_series[frame_num] - zoom = args.zoom_series[frame_num] - translation_x = args.translation_x_series[frame_num] - translation_y = args.translation_y_series[frame_num] - print( - f'angle: {angle}', - f'zoom: {zoom}', - f'translation_x: {translation_x}', - f'translation_y: {translation_y}', - ) - - if frame_num > 0: - seed += 1 - if args.resume_run and frame_num == args.start_frame: - img_0 = cv2.imread( - args.batchFolder+f"/{args.batch_name}({batchNum})_{args.start_frame-1:04}.png") - else: - img_0 = cv2.imread('prevFrame.png') - center = (1*img_0.shape[1]//2, 1*img_0.shape[0]//2) - trans_mat = np.float32( - [[1, 0, translation_x], - [0, 1, translation_y]] - ) - rot_mat = cv2.getRotationMatrix2D(center, angle, zoom) - trans_mat = np.vstack([trans_mat, [0, 0, 1]]) - rot_mat = np.vstack([rot_mat, [0, 0, 1]]) - transformation_matrix = np.matmul(rot_mat, trans_mat) - img_0 = cv2.warpPerspective( - img_0, - transformation_matrix, - (img_0.shape[1], img_0.shape[0]), - borderMode=cv2.BORDER_WRAP - ) - - cv2.imwrite('prevFrameScaled.png', img_0) - init_image = 'prevFrameScaled.png' - init_scale = args.frames_scale - skip_steps = args.calc_frames_skip_steps - - if args.animation_mode == "3D": - if frame_num > 0: - seed += 1 - if args.resume_run and frame_num == args.start_frame: - img_filepath = args.batchFolder + \ - f"/{args.batch_name}({batchNum})_{args.start_frame-1:04}.png" - if args.turbo_mode and frame_num > args.turbo_preroll: - shutil.copyfile(img_filepath, 'oldFrameScaled.png') - else: - img_filepath = 'prevFrame.png' - - next_step_pil = do_3d_step( - args, img_filepath, frame_num, midas_model, midas_transform) - next_step_pil.save('prevFrameScaled.png') - - # Turbo mode - skip some diffusions, use 3d morph for clarity and to save time - if args.turbo_mode: - if frame_num == args.turbo_preroll: # start tracking oldframe - # stash for later blending - next_step_pil.save('oldFrameScaled.png') - elif frame_num > args.turbo_preroll: - # set up 2 warped image sequences, old & new, to blend toward new diff image - old_frame = do_3d_step( - args, 'oldFrameScaled.png', frame_num, midas_model, midas_transform) - old_frame.save('oldFrameScaled.png') - if frame_num % int(args.turbo_steps) != 0: - print( - 'turbo skip this frame: skipping clip diffusion steps') - filename = f'{args.batch_name}({batchNum})_{frame_num:04}.png' - blend_factor = ( - (frame_num % int(args.turbo_steps))+1)/int(args.turbo_steps) - print( - 'turbo skip this frame: skipping clip diffusion steps and saving blended frame') - # this is already updated.. - newWarpedImg = cv2.imread('prevFrameScaled.png') - oldWarpedImg = cv2.imread('oldFrameScaled.png') - blendedImage = cv2.addWeighted( - newWarpedImg, blend_factor, oldWarpedImg, 1-blend_factor, 0.0) - cv2.imwrite( - f'{args.batchFolder}/{filename}', blendedImage) - # save it also as prev_frame to feed next iteration - next_step_pil.save(f'{img_filepath}') - if args.vr_mode: - generate_eye_views( - TRANSLATION_SCALE, args.batchFolder, filename, frame_num, midas_model, midas_transform) - continue - else: - # if not a skip frame, will run diffusion and need to blend. - oldWarpedImg = cv2.imread('prevFrameScaled.png') - # swap in for blending later - cv2.imwrite(f'oldFrameScaled.png', oldWarpedImg) - print('clip/diff this frame - generate clip diff image') - - init_image = 'prevFrameScaled.png' - init_scale = args.frames_scale - skip_steps = args.calc_frames_skip_steps - - if args.animation_mode == "Video Input": - init_scale = args.video_init_frames_scale + cv2.imwrite('prevFrameScaled.png', img_0) + init_image = 'prevFrameScaled.png' + init_scale = args.frames_scale skip_steps = args.calc_frames_skip_steps - if not args.video_init_seed_continuity: - seed += 1 - if args.video_init_flow_warp: - if frame_num == 0: - skip_steps = args.video_init_skip_steps - init_image = f'{args.videoFramesFolder}/{frame_num+1:04}.jpg' - if frame_num > 0: - prev = PIL.Image.open( - args.batchFolder+f"/{args.batch_name}({batchNum})_{frame_num-1:04}.png") - - frame1_path = f'{args.videoFramesFolder}/{frame_num:04}.jpg' - frame2 = PIL.Image.open( - f'{args.videoFramesFolder}/{frame_num+1:04}.jpg') - flo_path = f"/{args.flo_folder}/{frame1_path.split('/')[-1]}.npy" - - init_image = 'warped.png' - print(args.video_init_flow_blend) - weights_path = None - if args.video_init_check_consistency: - # TBD - pass - - import video_input - video_input.warp(prev, frame2, flo_path, blend=args.video_init_flow_blend, - weights_path=weights_path).save(init_image) + if args.animation_mode == "3D": + if frame_num > 0: + args.seed += 1 + if args.resume_run and frame_num == args.start_frame: + img_filepath = args.batchFolder + \ + f"/{args.batch_name}({batchNum})_{args.start_frame-1:04}.png" + if args.turbo_mode and frame_num > args.turbo_preroll: + shutil.copyfile(img_filepath, 'oldFrameScaled.png') else: + img_filepath = 'prevFrame.png' + + next_step_pil = do_3d_step( + args, img_filepath, frame_num, midas_model, midas_transform) + next_step_pil.save('prevFrameScaled.png') + + # Turbo mode - skip some diffusions, use 3d morph for clarity and to save time + if args.turbo_mode: + if frame_num == args.turbo_preroll: # start tracking oldframe + # stash for later blending + next_step_pil.save('oldFrameScaled.png') + elif frame_num > args.turbo_preroll: + # set up 2 warped image sequences, old & new, to blend toward new diff image + old_frame = do_3d_step( + args, 'oldFrameScaled.png', frame_num, midas_model, midas_transform) + old_frame.save('oldFrameScaled.png') + if frame_num % int(args.turbo_steps) != 0: + print( + 'turbo skip this frame: skipping clip diffusion steps') + filename = f'{args.batch_name}({batchNum})_{frame_num:04}.png' + blend_factor = ( + (frame_num % int(args.turbo_steps))+1)/int(args.turbo_steps) + print( + 'turbo skip this frame: skipping clip diffusion steps and saving blended frame') + # this is already updated.. + newWarpedImg = cv2.imread('prevFrameScaled.png') + oldWarpedImg = cv2.imread('oldFrameScaled.png') + blendedImage = cv2.addWeighted( + newWarpedImg, blend_factor, oldWarpedImg, 1-blend_factor, 0.0) + cv2.imwrite( + f'{args.batchFolder}/{filename}', blendedImage) + # save it also as prev_frame to feed next iteration + next_step_pil.save(f'{img_filepath}') + if args.vr_mode: + generate_eye_views( + TRANSLATION_SCALE, args.batchFolder, filename, frame_num, midas_model, midas_transform) + return + else: + # if not a skip frame, will run diffusion and need to blend. + oldWarpedImg = cv2.imread('prevFrameScaled.png') + # swap in for blending later + cv2.imwrite(f'oldFrameScaled.png', oldWarpedImg) + print('clip/diff this frame - generate clip diff image') + + init_image = 'prevFrameScaled.png' + init_scale = args.frames_scale + skip_steps = args.calc_frames_skip_steps + + if args.animation_mode == "Video Input": + init_scale = args.video_init_frames_scale + skip_steps = args.calc_frames_skip_steps + if not args.video_init_seed_continuity: + args.seed += 1 + if args.video_init_flow_warp: + if frame_num == 0: + skip_steps = args.video_init_skip_steps init_image = f'{args.videoFramesFolder}/{frame_num+1:04}.jpg' + if frame_num > 0: + prev = PIL.Image.open( + args.batchFolder+f"/{args.batch_name}({batchNum})_{frame_num-1:04}.png") - loss_values = [] + frame1_path = f'{args.videoFramesFolder}/{frame_num:04}.jpg' + frame2 = PIL.Image.open( + f'{args.videoFramesFolder}/{frame_num+1:04}.jpg') + flo_path = f"/{args.flo_folder}/{frame1_path.split('/')[-1]}.npy" - if seed is not None: - np.random.seed(seed) - random.seed(seed) - torch.manual_seed(seed) - torch.cuda.manual_seed_all(seed) - torch.backends.cudnn.deterministic = True + init_image = 'warped.png' + print(args.video_init_flow_blend) + weights_path = None + if args.video_init_check_consistency: + # TBD + pass - target_embeds, weights = [], [] + import video_input + video_input.warp(prev, frame2, flo_path, blend=args.video_init_flow_blend, + weights_path=weights_path).save(init_image) - if args.prompts_series is not None and frame_num >= len(args.prompts_series): - frame_prompt = args.prompts_series[-1] - elif args.prompts_series is not None: - frame_prompt = args.prompts_series[frame_num] else: - frame_prompt = [] + init_image = f'{args.videoFramesFolder}/{frame_num+1:04}.jpg' - print(args.image_prompts_series) - if args.image_prompts_series is not None and frame_num >= len(args.image_prompts_series): - image_prompt = args.image_prompts_series[-1] - elif args.image_prompts_series is not None: - image_prompt = args.image_prompts_series[frame_num] - else: - image_prompt = [] + loss_values = [] - device = comfy.model_management.get_torch_device() + if args.seed is not None: + np.random.seed(args.seed) + random.seed(args.seed) + torch.manual_seed(args.seed) + torch.cuda.manual_seed_all(args.seed) + torch.backends.cudnn.deterministic = True - print(f'Frame {frame_num} Prompt: {frame_prompt}') + target_embeds, weights = [], [] - # from .CLIP import clip as openai_clip - # clip_model = openai_clip.load('ViT-B/32', jit=False)[0].eval().requires_grad_(False).to(device) - # clip_vision = clip_model + if args.prompts_series is not None and frame_num >= len(args.prompts_series): + frame_prompt = args.prompts_series[-1] + elif args.prompts_series is not None: + frame_prompt = args.prompts_series[frame_num] + else: + frame_prompt = [] - clip_models = [clip_model] # TODO!!!!!!!!!!!!!!!!!!!!! + print(args.image_prompts_series) + if args.image_prompts_series is not None and frame_num >= len(args.image_prompts_series): + image_prompt = args.image_prompts_series[-1] + elif args.image_prompts_series is not None: + image_prompt = args.image_prompts_series[frame_num] + else: + image_prompt = [] - model_stats = [] - for clip_model in clip_models: - cutn = 16 - model_stat = {"clip_model": None, "target_embeds": [], - "make_cutouts": None, "weights": []} - model_stat["clip_model"] = clip_model - model_stat["clip_vision_model"] = clip_vision + device = comfy.model_management.get_torch_device() - for prompt in frame_prompt: - prompt = ", ".join(prompt) - txt, weight = disco_utils.parse_prompt(prompt) - txt = encode_text(clip_model, prompt).to(device) - # txt = clip_model.encode(prompt).float() - # txt = clip_model.encode_text(openai_clip.tokenize(prompt).to(device)).float() + if isinstance(clip_vision, ClipVisionModel): + clip_vision.model.to(device) # Gets loaded to CPU by comfy, move to GPU + print(f'Frame {frame_num} Prompt: {frame_prompt}') + + + clip_models = [clip_model] + + model_stats = [] + for clip_model in clip_models: + cutn = 16 + model_stat = {"clip_model": None, "target_embeds": [], + "make_cutouts": None, "weights": []} + model_stat["clip_model"] = clip_model + model_stat["clip_vision_model"] = clip_vision + + for prompt in frame_prompt: + prompt = ", ".join(prompt) + txt, weight = disco_utils.parse_prompt(prompt) + txt = encode_text(clip_model, prompt).to(device) + + if args.fuzzy_prompt: + for i in range(25): + model_stat["target_embeds"].append( + (txt + torch.randn(txt.shape).cuda() * args.rand_mag).clamp(0, 1)) + model_stat["weights"].append(weight) + else: + model_stat["target_embeds"].append(txt) + model_stat["weights"].append(weight) + + if image_prompt: + input_res = get_input_resolution(clip_vision) + model_stat["make_cutouts"] = MakeCutouts(input_res, cutn, skip_augs=args.skip_augs) + for prompt in image_prompt: + path, weight = disco_utils.parse_prompt(prompt) + img = Image.open(disco_utils.fetch(path)).convert('RGB') + img = TF.resize( + img, min(args.side_x, args.side_y, *img.size), T.InterpolationMode.LANCZOS) + batch = model_stat["make_cutouts"](TF.to_tensor( + img).to(device).unsqueeze(0).mul(2).sub(1)) + embed = encode_images(clip_vision, disco_utils.normalize(batch)) if args.fuzzy_prompt: for i in range(25): model_stat["target_embeds"].append( - (txt + torch.randn(txt.shape).cuda() * args.rand_mag).clamp(0, 1)) - model_stat["weights"].append(weight) + (embed + torch.randn(embed.shape).cuda() * args.rand_mag).clamp(0, 1)) + weights.extend([weight / cutn] * cutn) else: - model_stat["target_embeds"].append(txt) - model_stat["weights"].append(weight) + model_stat["target_embeds"].append(embed) + model_stat["weights"].extend([weight / cutn] * cutn) - if image_prompt: - model_stat["make_cutouts"] = MakeCutouts( - clip_vision.model.config.image_size, cutn, skip_augs=args.skip_augs) - for prompt in image_prompt: - path, weight = disco_utils.parse_prompt(prompt) - img = Image.open(disco_utils.fetch(path)).convert('RGB') - img = TF.resize( - img, min(args.side_x, args.side_y, *img.size), T.InterpolationMode.LANCZOS) - batch = model_stat["make_cutouts"](TF.to_tensor( - img).to(device).unsqueeze(0).mul(2).sub(1)) - embed = encode_images(clip_vision, disco_utils.normalize(batch)) - if args.fuzzy_prompt: - for i in range(25): - model_stat["target_embeds"].append( - (embed + torch.randn(embed.shape).cuda() * args.rand_mag).clamp(0, 1)) - weights.extend([weight / cutn] * cutn) - else: - model_stat["target_embeds"].append(embed) - model_stat["weights"].extend([weight / cutn] * cutn) + model_stat["target_embeds"] = torch.cat( + model_stat["target_embeds"]) + model_stat["weights"] = torch.tensor( + model_stat["weights"], device=device) + if model_stat["weights"].sum().abs() < 1e-3: + raise RuntimeError('The weights must not sum to 0.') + model_stat["weights"] /= model_stat["weights"].sum().abs() + model_stats.append(model_stat) - model_stat["target_embeds"] = torch.cat( - model_stat["target_embeds"]) - model_stat["weights"] = torch.tensor( - model_stat["weights"], device=device) - if model_stat["weights"].sum().abs() < 1e-3: - raise RuntimeError('The weights must not sum to 0.') - model_stat["weights"] /= model_stat["weights"].sum().abs() - model_stats.append(model_stat) + init = None + if init_image is not None: + init = Image.open(disco_utils.fetch(init_image)).convert('RGB') + init = init.resize((args.side_x, args.side_y), Image.LANCZOS) + init = TF.to_tensor(init).to(device).unsqueeze(0).mul(2).sub(1) - init = None - if init_image is not None: - init = Image.open(disco_utils.fetch(init_image)).convert('RGB') - init = init.resize((args.side_x, args.side_y), Image.LANCZOS) - init = TF.to_tensor(init).to(device).unsqueeze(0).mul(2).sub(1) + if args.perlin_init: + init = disco_utils.regen_perlin(args.perlin_mode, args.batch_size) + + cur_t = None + + def cond_fn(x, t, y=None): + with torch.enable_grad(): + x_is_NaN = False + x = x.detach().requires_grad_() + n = x.shape[0] + if args.MS.use_secondary_model is True: + alpha = torch.tensor( + diffusion.sqrt_alphas_cumprod[cur_t], device=device, dtype=torch.float32) + sigma = torch.tensor( + diffusion.sqrt_one_minus_alphas_cumprod[cur_t], device=device, dtype=torch.float32) + cosine_t = disco_utils.alpha_sigma_to_t(alpha, sigma) + out = args.MS.secondary_model( + x, cosine_t[None].repeat([n])).pred + fac = diffusion.sqrt_one_minus_alphas_cumprod[cur_t] + x_in = out * fac + x * (1 - fac) + x_in_grad = torch.zeros_like(x_in) + else: + my_t = torch.ones([n], device=device, + dtype=torch.long) * cur_t + out = diffusion.p_mean_variance( + model, x, my_t, clip_denoised=False, model_kwargs={'y': y}) + fac = diffusion.sqrt_one_minus_alphas_cumprod[cur_t] + x_in = out['pred_xstart'] * fac + x * (1 - fac) + x_in_grad = torch.zeros_like(x_in) + for model_stat in model_stats: + for i in range(args.cutn_batches): + # errors on last step without +1, need to find source + t_int = int(t.item())+1 + input_resolution = get_input_resolution(model_stat["clip_vision_model"]) + + cuts = MakeCutoutsDango(animation_mode=args.animation_mode, + skip_augs=args.skip_augs, + cut_size=input_resolution, + Overview=args.cut_overview[1000-t_int], + InnerCrop=args.cut_innercut[1000-t_int], + IC_Size_Pow=args.cut_ic_pow[1000-t_int], + IC_Grey_P=args.cut_icgray_p[1000-t_int] + ) + clip_in = disco_utils.normalize( + cuts(x_in.add(1).div(2))) + image_embeds = encode_images(model_stat["clip_vision_model"], clip_in) + dists = disco_utils.spherical_dist_loss(image_embeds.unsqueeze( + 1), model_stat["target_embeds"].unsqueeze(0)) + dists = dists.view( + [args.cut_overview[1000-t_int]+args.cut_innercut[1000-t_int], n, -1]) + losses = dists.mul( + model_stat["weights"]).sum(2).mean(0) + # log loss, probably shouldn't do per cutn_batch + loss_values.append(losses.sum().item()) + grads = torch.autograd.grad(losses.sum() * args.clip_guidance_scale, x_in) + x_in_grad += grads[0] / args.cutn_batches + tv_losses = disco_utils.tv_loss(x_in) + if args.MS.use_secondary_model is True: + range_losses = disco_utils.range_loss(out) + else: + range_losses = disco_utils.range_loss(out['pred_xstart']) + sat_losses = torch.abs(x_in - x_in.clamp(min=-1, max=1)).mean() + loss = tv_losses.sum() * args.tv_scale + range_losses.sum() * \ + args.range_scale + sat_losses.sum() * args.sat_scale + if init is not None and init_scale: + init_losses = args.MS.lpips_model(x_in, init) + loss = loss + init_losses.sum() * init_scale + x_in_grad += torch.autograd.grad(loss, x_in)[0] + if torch.isnan(x_in_grad).any() == False: + grad = -torch.autograd.grad(x_in, x, x_in_grad)[0] + else: + # print("NaN'd") + x_is_NaN = True + grad = torch.zeros_like(x) + if args.clamp_grad and x_is_NaN == False: + magnitude = grad.square().mean().sqrt() + # min=-0.02, min=-clamp_max, + return grad * magnitude.clamp(max=args.clamp_max) / magnitude + return grad + + if args.MS.diffusion_sampling_mode == 'ddim': + sample_fn = diffusion.ddim_sample_loop_progressive + else: + sample_fn = diffusion.plms_sample_loop_progressive + + results = [] + + # image_display = Output() + for i in range(args.n_batches): + # if args.animation_mode == 'None': + # display.clear_output(wait=True) + # batchBar = tqdm(range(args.n_batches), desc="Batches") + # batchBar.n = i + # batchBar.refresh() + # display.display(image_display) + gc.collect() + torch.cuda.empty_cache() + cur_t = diffusion.num_timesteps - skip_steps - 1 + total_steps = cur_t if args.perlin_init: - init = disco_utils.regen_perlin(args.perlin_mode, args.batch_size) + init = disco_utils.regen_perlin( + args.perlin_mode, args.batch_size, True) - cur_t = None - - def cond_fn(x, t, y=None): - with torch.enable_grad(): - x_is_NaN = False - x = x.detach().requires_grad_() - n = x.shape[0] - if args.MS.use_secondary_model is True: - alpha = torch.tensor( - diffusion.sqrt_alphas_cumprod[cur_t], device=device, dtype=torch.float32) - sigma = torch.tensor( - diffusion.sqrt_one_minus_alphas_cumprod[cur_t], device=device, dtype=torch.float32) - cosine_t = disco_utils.alpha_sigma_to_t(alpha, sigma) - out = args.MS.secondary_model( - x, cosine_t[None].repeat([n])).pred - fac = diffusion.sqrt_one_minus_alphas_cumprod[cur_t] - x_in = out * fac + x * (1 - fac) - x_in_grad = torch.zeros_like(x_in) - else: - my_t = torch.ones([n], device=device, - dtype=torch.long) * cur_t - out = diffusion.p_mean_variance( - model, x, my_t, clip_denoised=False, model_kwargs={'y': y}) - fac = diffusion.sqrt_one_minus_alphas_cumprod[cur_t] - x_in = out['pred_xstart'] * fac + x * (1 - fac) - x_in_grad = torch.zeros_like(x_in) - for model_stat in model_stats: - for i in range(args.cutn_batches): - # errors on last step without +1, need to find source - t_int = int(t.item())+1 - # when using SLIP Base model the dimensions need to be hard coded to avoid AttributeError: 'VisionTransformer' object has no attribute 'input_resolution' - try: - input_resolution = model_stat["clip_vision_model"].model.config.image_size - except Exception as err: - print("Couldn't find clip vision image size! " + str(err)) - input_resolution = 224 - - cuts = MakeCutoutsDango(animation_mode=args.animation_mode, - skip_augs=args.skip_augs, - cut_size=input_resolution, - Overview=args.cut_overview[1000-t_int], - InnerCrop=args.cut_innercut[1000-t_int], - IC_Size_Pow=args.cut_ic_pow[1000-t_int], - IC_Grey_P=args.cut_icgray_p[1000-t_int] - ) - clip_in = disco_utils.normalize( - cuts(x_in.add(1).div(2))) - image_embeds = encode_images(model_stat["clip_vision_model"], clip_in).to(device) - # image_embeds = model_stat["clip_model"].encode_image(clip_in).float() - print(image_embeds.shape) - print(model_stat["target_embeds"].shape) - print(image_embeds.unsqueeze(1).shape) - print(model_stat["target_embeds"].unsqueeze(0).shape) - dists = disco_utils.spherical_dist_loss(image_embeds.unsqueeze( - 1), model_stat["target_embeds"].unsqueeze(0)) - dists = dists.view( - [args.cut_overview[1000-t_int]+args.cut_innercut[1000-t_int], n, -1]) - losses = dists.mul( - model_stat["weights"]).sum(2).mean(0) - # log loss, probably shouldn't do per cutn_batch - loss_values.append(losses.sum().item()) - x_in_grad += torch.autograd.grad(losses.sum() * args.clip_guidance_scale, x_in)[ - 0] / args.cutn_batches - tv_losses = disco_utils.tv_loss(x_in) - if args.MS.use_secondary_model is True: - range_losses = disco_utils.range_loss(out) - else: - range_losses = disco_utils.range_loss(out['pred_xstart']) - sat_losses = torch.abs(x_in - x_in.clamp(min=-1, max=1)).mean() - loss = tv_losses.sum() * args.tv_scale + range_losses.sum() * \ - args.range_scale + sat_losses.sum() * args.sat_scale - if init is not None and init_scale: - init_losses = args.MS.lpips_model(x_in, init) - loss = loss + init_losses.sum() * init_scale - x_in_grad += torch.autograd.grad(loss, x_in)[0] - if torch.isnan(x_in_grad).any() == False: - grad = -torch.autograd.grad(x_in, x, x_in_grad)[0] - else: - # print("NaN'd") - x_is_NaN = True - grad = torch.zeros_like(x) - if args.clamp_grad and x_is_NaN == False: - magnitude = grad.square().mean().sqrt() - # min=-0.02, min=-clamp_max, - return grad * magnitude.clamp(max=args.clamp_max) / magnitude - return grad + symmetry_transformation_fn = id + if args.use_horizontal_symmetry: + symmetry_transformation_fn = horiz_symmetry + if args.use_vertical_symmetry: + symmetry_transformation_fn = vert_symmetry if args.MS.diffusion_sampling_mode == 'ddim': - sample_fn = diffusion.ddim_sample_loop_progressive + samples = sample_fn( + model, + (args.batch_size, 3, args.side_y, args.side_x), + clip_denoised=args.clip_denoised, + model_kwargs={}, + cond_fn=cond_fn, + progress=True, + skip_timesteps=skip_steps, + init_image=init, + randomize_class=args.randomize_class, + eta=args.eta, + transformation_fn=symmetry_transformation_fn, + transformation_percent=args.transformation_percent + ) else: - sample_fn = diffusion.plms_sample_loop_progressive + samples = sample_fn( + model, + (args.batch_size, 3, args.side_y, args.side_x), + clip_denoised=args.clip_denoised, + model_kwargs={}, + cond_fn=cond_fn, + progress=True, + skip_timesteps=skip_steps, + init_image=init, + randomize_class=args.randomize_class, + order=2, + ) - # image_display = Output() - for i in range(args.n_batches): - # if args.animation_mode == 'None': - # display.clear_output(wait=True) - # batchBar = tqdm(range(args.n_batches), desc="Batches") - # batchBar.n = i - # batchBar.refresh() - print(f"+++ Batch {i} +++") - # display.display(image_display) - gc.collect() - torch.cuda.empty_cache() - cur_t = diffusion.num_timesteps - skip_steps - 1 - total_steps = cur_t - - if args.perlin_init: - init = disco_utils.regen_perlin( - args.perlin_mode, args.batch_size, True) - - symmetry_transformation_fn = id - if args.use_horizontal_symmetry: - symmetry_transformation_fn = horiz_symmetry - if args.use_vertical_symmetry: - symmetry_transformation_fn = vert_symmetry - - if args.MS.diffusion_sampling_mode == 'ddim': - samples = sample_fn( - model, - (args.batch_size, 3, args.side_y, args.side_x), - clip_denoised=args.clip_denoised, - model_kwargs={}, - cond_fn=cond_fn, - progress=True, - skip_timesteps=skip_steps, - init_image=init, - randomize_class=args.randomize_class, - eta=args.eta, - transformation_fn=symmetry_transformation_fn, - transformation_percent=args.transformation_percent - ) - else: - samples = sample_fn( - model, - (args.batch_size, 3, args.side_y, args.side_x), - clip_denoised=args.clip_denoised, - model_kwargs={}, - cond_fn=cond_fn, - progress=True, - skip_timesteps=skip_steps, - init_image=init, - randomize_class=args.randomize_class, - order=2, - ) - - # with run_display: - # display.clear_output(wait=True) - for j, sample in enumerate(samples): - pbar.update_absolute(j, diffusion.num_timesteps - skip_steps) - cur_t -= 1 - intermediateStep = False - if args.steps_per_checkpoint is not None: - if j % args.steps_per_checkpoint == 0 and j > 0: - intermediateStep = True - elif j in args.intermediate_saves: + # with run_display: + # display.clear_output(wait=True) + for j, sample in enumerate(samples): + pbar.update_absolute(j, diffusion.num_timesteps - skip_steps) + cur_t -= 1 + intermediateStep = False + if args.steps_per_checkpoint is not None: + if j % args.steps_per_checkpoint == 0 and j > 0: intermediateStep = True - # with image_display: - if j % args.display_rate == 0 or cur_t == -1 or intermediateStep == True: - for k, image in enumerate(sample['pred_xstart']): - # tqdm.write(f'Batch {i}, step {j}, output {k}:') - datetime.now().strftime('%y%m%d-%H%M%S_%f') - percent = math.ceil(j/total_steps*100) - if args.n_batches > 0: - # if intermediates are saved to the subfolder, don't append a step or percentage to the name - if cur_t == -1 and args.intermediates_in_subfolder is True: - save_num = f'{frame_num:04}' if args.animation_mode != "None" else i - filename = f'{args.batch_name}({batchNum})_{save_num}.png' - else: - # If we're working with percentages, append it - if args.steps_per_checkpoint is not None: - filename = f'{args.batch_name}({batchNum})_{i:04}-{percent:02}%.png' - # Or else, iIf we're working with specific steps, append those - else: - filename = f'{args.batch_name}({batchNum})_{i:04}-{j:03}.png' - image = TF.to_pil_image( - image.add(1).div(2).clamp(0, 1)) - if j % args.display_rate == 0 or cur_t == -1: - image.save('progress.png') - # display.clear_output(wait=True) - # display.display(display.Image('progress.png')) - if args.steps_per_checkpoint is not None: - if j % args.steps_per_checkpoint == 0 and j > 0: - if args.intermediates_in_subfolder is True: - image.save( - f'{args.partialFolder}/{filename}') - else: - image.save( - f'{args.batchFolder}/{filename}') + elif j in args.intermediate_saves: + intermediateStep = True + # with image_display: + if j % args.display_rate == 0 or cur_t == -1 or intermediateStep == True: + for k, image in enumerate(sample['pred_xstart']): + # tqdm.write(f'Batch {i}, step {j}, output {k}:') + datetime.now().strftime('%y%m%d-%H%M%S_%f') + percent = math.ceil(j/total_steps*100) + if args.n_batches > 0: + # if intermediates are saved to the subfolder, don't append a step or percentage to the name + if cur_t == -1 and args.intermediates_in_subfolder is True: + save_num = f'{frame_num:04}' if args.animation_mode != "None" else i + filename = f'{args.batch_name}({batchNum})_{save_num}.png' else: - if j in args.intermediate_saves: - if args.intermediates_in_subfolder is True: - image.save( - f'{args.partialFolder}/{filename}') - else: - image.save( - f'{args.batchFolder}/{filename}') - if cur_t == -1: - # if frame_num == 0: - # save_settings() - if args.animation_mode != "None": - image.save('prevFrame.png') - image.save(f'{args.batchFolder}/{filename}') - if args.animation_mode == "3D": - # If turbo, save a blended image - if args.turbo_mode and frame_num > 0: - # Mix new image with prevFrameScaled - blend_factor = (1)/int(args.turbo_steps) - # This is already updated.. - newFrame = cv2.imread('prevFrame.png') - prev_frame_warped = cv2.imread( - 'prevFrameScaled.png') - blendedImage = cv2.addWeighted( - newFrame, blend_factor, prev_frame_warped, (1-blend_factor), 0.0) - cv2.imwrite( - f'{args.batchFolder}/{filename}', blendedImage) - else: - image.save( - f'{args.batchFolder}/{filename}') + # If we're working with percentages, append it + if args.steps_per_checkpoint is not None: + filename = f'{args.batch_name}({batchNum})_{i:04}-{percent:02}%.png' + # Or else, iIf we're working with specific steps, append those + else: + filename = f'{args.batch_name}({batchNum})_{i:04}-{j:03}.png' + save_image(image, j, cur_t, filename, frame_num, midas_model, midas_transform, args) - if args.vr_mode: - generate_eye_views( - args, TRANSLATION_SCALE, args.batchFolder, filename, frame_num, midas_model, midas_transform) + results.append(image) - # if frame_num != args.max_frames-1: - # display.clear_output() + # plt.plot(np.array(loss_values), 'r') + return torch.stack(results) - # plt.plot(np.array(loss_values), 'r') +def save_image(image, j, cur_t, filename, frame_num, midas_model, midas_transform, args): + image = TF.to_pil_image( + image.add(1).div(2).clamp(0, 1)) + # if j % args.display_rate == 0 or cur_t == -1: + # image.save('progress.png') + # display.clear_output(wait=True) + # display.display(display.Image('progress.png')) + if args.steps_per_checkpoint is not None: + if j % args.steps_per_checkpoint == 0 and j > 0: + if args.intermediates_in_subfolder is True: + image.save( + f'{args.partialFolder}/{filename}') + else: + image.save( + f'{args.batchFolder}/{filename}') + else: + if j in args.intermediate_saves: + if args.intermediates_in_subfolder is True: + image.save( + f'{args.partialFolder}/{filename}') + else: + image.save( + f'{args.batchFolder}/{filename}') + if cur_t == -1: + # if frame_num == 0: + # save_settings() + if args.animation_mode != "None": + image.save('prevFrame.png') + image.save(f'{args.batchFolder}/{filename}') + if args.animation_mode == "3D": + # If turbo, save a blended image + if args.turbo_mode and frame_num > 0: + # Mix new image with prevFrameScaled + blend_factor = (1)/int(args.turbo_steps) + # This is already updated.. + newFrame = cv2.imread('prevFrame.png') + prev_frame_warped = cv2.imread( + 'prevFrameScaled.png') + blendedImage = cv2.addWeighted( + newFrame, blend_factor, prev_frame_warped, (1-blend_factor), 0.0) + cv2.imwrite( + f'{args.batchFolder}/{filename}', blendedImage) + else: + image.save( + f'{args.batchFolder}/{filename}') + + if args.vr_mode: + generate_eye_views( + args, TRANSLATION_SCALE, args.batchFolder, filename, frame_num, midas_model, midas_transform) + + # if frame_num != args.max_frames-1: + # display.clear_output() def generate_eye_views(args, trans_scale, batchFolder, filename, frame_num, midas_model, midas_transform): diff --git a/model_settings.py b/model_settings.py index 89df8d9..e97843b 100644 --- a/model_settings.py +++ b/model_settings.py @@ -43,11 +43,10 @@ diff_model_map = { } class ModelSettings: - def __init__(self): - self.root_path = os.getcwd() - self.model_path = f'{self.root_path}/models' + def __init__(self, model_name, model_path): + self.model_path = model_path #@markdown ####**Models Settings (note: For pixel art, the best is pixelartdiffusion_expanded):** - self.diffusion_model = "512x512_diffusion_uncond_finetune_008100" #@param ["256x256_diffusion_uncond", "512x512_diffusion_uncond_finetune_008100", "portrait_generator_v001", "pixelartdiffusion_expanded", "pixel_art_diffusion_hard_256", "pixel_art_diffusion_soft_256", "pixelartdiffusion4k", "watercolordiffusion_2", "watercolordiffusion", "PulpSciFiDiffusion", "custom"] + self.diffusion_model = model_name #@param ["256x256_diffusion_uncond", "512x512_diffusion_uncond_finetune_008100", "portrait_generator_v001", "pixelartdiffusion_expanded", "pixel_art_diffusion_hard_256", "pixel_art_diffusion_soft_256", "pixelartdiffusion4k", "watercolordiffusion_2", "watercolordiffusion", "PulpSciFiDiffusion", "custom"] self.use_secondary_model = True #@param {type: 'boolean'} self.diffusion_sampling_mode = 'ddim' #@param ['plms','ddim'] @@ -123,7 +122,7 @@ class ModelSettings: print(f'{diffusion_model_name} model download from {model_uri} failed. Will try any fallback uri.') print(f'{diffusion_model_name} download failed.') - def setup(self, S): + def setup(self, useCPU): # Download the diffusion model(s) self.download_model(self.diffusion_model) if self.use_secondary_model: @@ -145,7 +144,7 @@ class ModelSettings: 'num_res_blocks': 2, 'resblock_updown': True, 'use_checkpoint': self.use_checkpoint, - 'use_fp16': not S.useCPU, + 'use_fp16': not useCPU, 'use_scale_shift_norm': True, }) elif self.diffusion_model == '256x256_diffusion_uncond': @@ -163,7 +162,7 @@ class ModelSettings: 'num_res_blocks': 2, 'resblock_updown': True, 'use_checkpoint': self.use_checkpoint, - 'use_fp16': not S.useCPU, + 'use_fp16': not useCPU, 'use_scale_shift_norm': True, }) elif self.diffusion_model == 'portrait_generator_v001': @@ -235,5 +234,3 @@ class ModelSettings: #if self.RN101_quickgelu_yfcc15m: clip_models.append(open_clip.create_model('RN101-quickgelu', pretrained='yfcc15m').eval().requires_grad_(False).to(device)) self.lpips_model = lpips.LPIPS(net='vgg').to(device) - - S.MS = self diff --git a/nodes.py b/nodes.py index 0a9f4ba..15c10f0 100644 --- a/nodes.py +++ b/nodes.py @@ -1,5 +1,8 @@ import os.path import comfy.model_management +from comfy.cli_args import args +import folder_paths +from pprint import pp NODE_FILE = os.path.abspath(__file__) DISCO_DIFFUSION_ROOT = os.path.dirname(NODE_FILE) @@ -10,72 +13,215 @@ sys.path.append(os.path.join(DISCO_DIFFUSION_ROOT, "MiDaS")) sys.path.append(os.path.join(DISCO_DIFFUSION_ROOT, "ResizeRight")) sys.path.append(os.path.join(DISCO_DIFFUSION_ROOT, "guided-diffusion")) sys.path.append(os.path.join(DISCO_DIFFUSION_ROOT, "RAFT/core")) +sys.path.append(os.path.join(DISCO_DIFFUSION_ROOT, "open_clip/src")) +import torch +from guided_diffusion.script_util import create_model_and_diffusion, model_and_diffusion_defaults from .settings import DiscoDiffusionSettings -from .model_settings import ModelSettings +from .model_settings import ModelSettings, diff_model_map from .diffuse import diffuse +from .CLIP import clip as openai_clip +import open_clip -# class DiscoDiffusionCLIPLoader: -# """ -# Loader for CLIP models compatible with Disco Diffusion (VIT-B) -# """ +OPENAI_CLIP_MODELS = ['ViT-B/32', 'ViT-B/16', 'ViT-L/14', 'ViT-L/14@336px', 'RN50', 'RN50x4', 'RN50x16', 'RN50x64', 'RN101'] +OPEN_CLIP_MODELS = [ + ['ViT-B-32', 'laion2b_e16'], + ['ViT-B-32', 'laion400m_e31'], + ['ViT-B-32', 'laion400m_e32'], + ['ViT-B-32-quickgelu', 'laion400m_e31'], + ['ViT-B-32-quickgelu', 'laion400m_e32'], + ['ViT-B-16', 'laion400m_e31'], + ['ViT-B-16', 'laion400m_e32'], + ['RN50', 'yfcc15m'], + ['RN50', 'cc12m'], + ['RN50-quickgelu', 'yfcc15m'], + ['RN50-quickgelu', 'cc12m'], + ['RN101', 'yfcc15m'], + ['RN101-quickgelu', 'yfcc15m'] +] -# @classmethod -# def INPUT_TYPES(s): -# return {"required": { "clip_model_name": (folder_paths.get_filename_list("style_models"), )}} -# RETURN_TYPES = () -# FUNCTION = "generate" +class OpenAICLIPLoader: + @classmethod + def INPUT_TYPES(s): + open_clip_models = ["_".join(m) for m in OPEN_CLIP_MODELS] + return {"required": {"model_name": (OPENAI_CLIP_MODELS + open_clip_models, { "default": "ViT-B/32" }) }} -# CATEGORY = "sampling" + # These are technically different model formats so don't use them with vanilla nodes! + RETURN_TYPES = ("CLIP", "CLIP_VISION") + FUNCTION = "load" -# OUTPUT_NODE = True + CATEGORY = "loaders" -# def __init__(self): -# self.settings = DiscoDiffusionSettings() -# self.model_settings = ModelSettings() -# self.settings.setup(self.model_settings) -# self.model_settings.setup(self.settings) + def __init__(self): + pass -# def generate(self, clip, clip_vision, text, seed): -# device = comfy.model_management.get_torch_device() -# diffuse(clip, clip_vision, self.settings, 0) -# return { "ui": { "images": {} } } + def load(self, model_name): + device = comfy.model_management.get_torch_device() + + if model_name in OPENAI_CLIP_MODELS: + clip_model = openai_clip.load(model_name, jit=False)[0] + else: + spl = model_name.split("_") + clip_model = open_clip.create_model(spl[0], pretrained=spl[1]) + + clip_model.eval().requires_grad_(False).to(device) + + return (clip_model, clip_model,) + + +GUIDED_DIFFUSION_MODELS = list(diff_model_map.keys()) + + +class GuidedDiffusionLoader: + @classmethod + def INPUT_TYPES(s): + return {"required": {"model_name": (GUIDED_DIFFUSION_MODELS, { "default": "512x512_diffusion_uncond_finetune_008100" }) }} + + # These are technically different model formats so don't use them with vanilla nodes! + RETURN_TYPES = ("GUIDED_DIFFUSION_MODEL",) + FUNCTION = "load" + + CATEGORY = "loaders" + + def __init__(self): + pass + + def load(self, model_name): + use_cpu = args.cpu + model_settings = ModelSettings(model_name, os.path.join(folder_paths.models_dir, "Disco-Diffusion")) + + model_settings.setup(use_cpu) + + return (model_settings,) + + +DEFAULT_PROMPT = """\ +# How to prompt: +# Each line is prefixed with the starting frame number of the prompt. +# More than one line with the same frame number concatenates the two prompts together. +# Each individual prompt can be no more than 77 characters long. +# Weights are parsed from the end of each prompt with "25:a fluffy fox:5" syntax +# Comments are written with the '#' character. Blank lines are ignored. + +0:A beautiful painting of a singular lighthouse, shining its light across a tumultuous sea of blood by greg rutkowski and thomas kinkade. Trending on artstation. +0:yellow color scheme +100:This set of prompts start at frame 100. +100:This prompt has weight five:5 +""".strip() class DiscoDiffusion: @classmethod def INPUT_TYPES(s): - return {"required": {"text": ("STRING", {"multiline": True}), + return {"required": {"text": ("STRING", {"default": DEFAULT_PROMPT, "multiline": True}), + "guided_diffusion": ("GUIDED_DIFFUSION_MODEL",), "clip": ("CLIP",), "clip_vision": ("CLIP_VISION",), + "steps": ("INT", {"default": 250, "min": 1, "max": 10000}), + "n_batches": ("INT", {"default": 1, "min": 1, "max": 16}), + # "max_frames": ("INT", {"default": 1, "min": 1, "max": 1000}), "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), }} - RETURN_TYPES = () + RETURN_TYPES = ("IMAGE",) + OUTPUT_IS_LIST = (True,) FUNCTION = "generate" CATEGORY = "sampling" - OUTPUT_NODE = True - def __init__(self): - self.settings = DiscoDiffusionSettings() - self.model_settings = ModelSettings() - self.settings.setup(self.model_settings) - self.model_settings.setup(self.settings) + pass - def generate(self, clip, clip_vision, text, seed): + def parse_prompts(self, text): + result = {} + for line in text.split('\n'): + line = line.split('#')[0].strip() + if line: + if ':' in line: + vals = line.split(':', 2) + key = vals[0] + value = vals[1] + weight = "1" + if len(vals) >= 3: + weight = vals[2] + + try: + key = int(key) + except ValueError: + weight = value + value = key + key = 0 + + if key in result: + result[key].append(value.strip() + ":" + weight) + else: + result[key] = [value.strip() + ":" + weight] + else: + if 0 in result: + result[0].append(line) + else: + result[0] = [line] + return result + + def load_model(self, model_settings, steps): device = comfy.model_management.get_torch_device() - diffuse(clip, clip_vision, self.settings, 0) - return { "ui": { "images": {} } } + + # Update Model Settings + timestep_respacing = f'ddim{steps}' + diffusion_steps = (1000//steps)*steps if steps < 1000 else steps + model_settings.model_config.update({ + 'timestep_respacing': timestep_respacing, + 'diffusion_steps': diffusion_steps, + }) + + model, diffusion = create_model_and_diffusion(**model_settings.model_config) + if model_settings.diffusion_model == 'custom': + model.load_state_dict(torch.load(model_settings.custom_path, map_location='cpu')) + else: + model.load_state_dict(torch.load(f'{model_settings.model_path}/{model_settings.get_model_filename(model_settings.diffusion_model)}', map_location='cpu')) + model.requires_grad_(False).eval().to(device) + + for name, param in model.named_parameters(): + if 'qkv' in name or 'norm' in name or 'proj' in name: + param.requires_grad_() + + if model_settings.model_config['use_fp16']: + model.convert_to_fp16() + + return model, diffusion + + def generate(self, text, guided_diffusion, clip, clip_vision, steps, n_batches, seed): + settings = DiscoDiffusionSettings() + settings.seed = seed + settings.steps = steps + settings.n_batches = n_batches + settings.max_frames = 1 + settings.text_prompts = self.parse_prompts(text) + + print("[Disco Diffusion] Parsed Prompts:") + pp(settings.text_prompts) + + settings.setup(guided_diffusion) + + # Have to defer loading the model until here since step count isn't + # known until now + model, diffusion = self.load_model(guided_diffusion, settings.steps) + + images = diffuse(model, diffusion, clip, clip_vision, settings, 0) + + return (images,) NODE_CLASS_MAPPINGS = { + "ComfyUI_OpenAICLIPLoader": OpenAICLIPLoader, + "ComfyUI_GuidedDiffusionLoader": GuidedDiffusionLoader, "ComfyUI_DiscoDiffusion": DiscoDiffusion, } NODE_DISPLAY_NAME_MAPPINGS = { + "ComfyUI_OpenAICLIPLoader": "OpenAI CLIP Loader", + "ComfyUI_GuidedDiffusionLoader": "Guided Diffusion Loader", "ComfyUI_DiscoDiffusion": "Disco Diffusion", } diff --git a/settings.py b/settings.py index 754b5d4..610b3a0 100644 --- a/settings.py +++ b/settings.py @@ -49,6 +49,7 @@ class DiscoDiffusionSettings: # @markdown ####**Basic Settings:** self.batch_name = 'TimeToDisco' # @param{type: 'string'} self.batch_size = 1 + self.n_batches = 1 # @param [25,50,100,150,250,500,1000]{type: 'raw', allow-input: true} self.steps = 250 self.width_height_for_512x512_models = [ @@ -112,7 +113,7 @@ class DiscoDiffusionSettings: # @markdown All rotations are provided in degrees. self.key_frames = True # @param {type:"boolean"} - self.max_frames = 100 # @param {type:"number"} + self.max_frames = 10 # @param {type:"number"} # Do not change, currently will not look good. param ['Linear','Quadratic','Cubic']{type:"string"} self.interp_spline = 'Linear' @@ -216,6 +217,7 @@ class DiscoDiffusionSettings: self.perlin_init = False # @param{type: 'boolean'} self.perlin_mode = 'mixed' # @param ['mixed', 'color', 'gray'] + self.seed = 0 self.set_seed = 'random_seed' # @param{type: 'string'} self.eta = 0.8 # @param{type: 'number'} self.clamp_grad = True # @param{type: 'boolean'}