Work!
This commit is contained in:
+6
-29
@@ -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
|
||||
|
||||
@@ -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):
|
||||
|
||||
+6
-9
@@ -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
|
||||
|
||||
@@ -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",
|
||||
}
|
||||
|
||||
+3
-1
@@ -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'}
|
||||
|
||||
Reference in New Issue
Block a user