Add looping option
This commit is contained in:
@@ -62,7 +62,7 @@ class DynamiCrafterModelLoader:
|
||||
model_config['params']['unet_config']['params']['use_checkpoint']=False
|
||||
self.model = instantiate_from_config(model_config)
|
||||
self.model = load_model_checkpoint(self.model, model_path)
|
||||
self.model.eval().to(dtype).to(device)
|
||||
self.model.eval().to(dtype)
|
||||
return (self.model,)
|
||||
|
||||
class DynamiCrafterI2V:
|
||||
@@ -83,7 +83,8 @@ class DynamiCrafterI2V:
|
||||
},
|
||||
"optional": {
|
||||
"image2": ("IMAGE",),
|
||||
"mask": ("MASK",),
|
||||
"mask": ("MASK",),
|
||||
"looping": ("BOOLEAN", {"default": False}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -92,7 +93,7 @@ class DynamiCrafterI2V:
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "DynamiCrafterWrapper"
|
||||
|
||||
def process(self, model, image, prompt, cfg, steps, eta, seed, fs, keep_model_loaded, frames, mask=None, image2=None):
|
||||
def process(self, model, image, prompt, cfg, steps, eta, seed, fs, keep_model_loaded, frames, mask=None, image2=None, looping=False):
|
||||
device = mm.get_torch_device()
|
||||
mm.unload_all_models()
|
||||
mm.soft_empty_cache()
|
||||
@@ -100,7 +101,7 @@ class DynamiCrafterI2V:
|
||||
torch.manual_seed(seed)
|
||||
dtype = model.dtype
|
||||
self.model = model
|
||||
|
||||
self.model.to(device)
|
||||
autocast_condition = (dtype != torch.float32) and not comfy.model_management.is_device_mps(device)
|
||||
with torch.autocast(comfy.model_management.get_autocast_device(device), dtype=dtype) if autocast_condition else nullcontext():
|
||||
image = image * 2 - 1
|
||||
@@ -131,6 +132,10 @@ class DynamiCrafterI2V:
|
||||
img_tensor_repeat[:,:,-1:,:,:] = z2
|
||||
else:
|
||||
img_tensor_repeat = repeat(z, 'b c t h w -> b c (repeat t) h w', repeat=frames)
|
||||
if looping:
|
||||
img_tensor_repeat = torch.zeros_like(img_tensor_repeat)
|
||||
img_tensor_repeat[:,:,:1,:,:] = z
|
||||
img_tensor_repeat[:,:,-1:,:,:] = z
|
||||
|
||||
self.model.first_stage_model.to('cpu')
|
||||
|
||||
@@ -213,7 +218,7 @@ class DynamiCrafterI2V:
|
||||
video = video.squeeze(0).permute(1, 2, 3, 0)
|
||||
|
||||
if not keep_model_loaded:
|
||||
self.model = None
|
||||
self.model.to('cpu')
|
||||
mm.soft_empty_cache()
|
||||
|
||||
last_image = video[-1].unsqueeze(0)
|
||||
@@ -249,7 +254,7 @@ class DynamiCrafterBatchInterpolation:
|
||||
torch.manual_seed(seed)
|
||||
dtype = model.dtype
|
||||
self.model = model
|
||||
|
||||
self.model.to(device)
|
||||
images = images * 2 - 1
|
||||
images = images.permute(0, 3, 1, 2).to(dtype).to(device)
|
||||
B, C, H, W = images.shape
|
||||
@@ -356,7 +361,7 @@ class DynamiCrafterBatchInterpolation:
|
||||
out.append(video)
|
||||
|
||||
if not keep_model_loaded:
|
||||
self.model = None
|
||||
self.model.to('cpu')
|
||||
mm.soft_empty_cache()
|
||||
out_video = torch.cat(out, dim=0)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user