diff --git a/__init__.py b/__init__.py index 14b5238..6a599e1 100644 --- a/__init__.py +++ b/__init__.py @@ -118,7 +118,11 @@ class LoadImageSequence: os.makedirs(s.input_dir) image_folder = [name for name in os.listdir(s.input_dir) if os.path.isdir(os.path.join(s.input_dir,name)) and len(os.listdir(os.path.join(s.input_dir,name))) != 0] return {"required": - {"image_sequence_folder": (sorted(image_folder), )} + {"image_sequence_folder": (sorted(image_folder), ), + "sample_start_idx": ("INT", {"default": 1, "min": 1, "max": 10000}), + "sample_frame_rate": ("INT", {"default": 1, "min": 1, "max": 10000}), + "n_sample_frames": ("INT", {"default": 1, "min": 1, "max": 10000}) + } } CATEGORY = "image" @@ -126,13 +130,14 @@ class LoadImageSequence: RETURN_TYPES = ("IMAGE", "MASK_SEQUENCE") FUNCTION = "load_image_sequence" - def load_image_sequence(self, image_sequence_folder): + def load_image_sequence(self, image_sequence_folder, sample_start_idx, sample_frame_rate, n_sample_frames): image_path = os.path.join(self.input_dir, image_sequence_folder) file_list = sorted(os.listdir(image_path), key=lambda s: sum(((s, int(n)) for s, n in re.findall(r'(\D+)(\d+)', 'a%s0' % s)), ())) sample_frames = [] sample_frames_mask = [] - for file in file_list: - i = Image.open(os.path.join(image_path, file)) + sample_index = list(range(sample_start_idx-1, len(file_list), sample_frame_rate))[:n_sample_frames] + for num in sample_index: + i = Image.open(os.path.join(image_path, file_list[num])) image = i.convert("RGB") image = np.array(image).astype(np.float32) / 255.0 image = torch.from_numpy(image)[None,] @@ -200,7 +205,11 @@ class LoadImageMaskSequence: image_folder = [name for name in os.listdir(s.input_dir) if os.path.isdir(os.path.join(s.input_dir, name)) and len(os.listdir(os.path.join(s.input_dir, name))) != 0] return {"required": {"image_sequence_folder": (sorted(image_folder), ), - "channel": (["alpha", "red", "green", "blue"], ),} + "channel": (["alpha", "red", "green", "blue"], ), + "sample_start_idx": ("INT", {"default": 1, "min": 1, "max": 10000}), + "sample_frame_rate": ("INT", {"default": 1, "min": 1, "max": 10000}), + "n_sample_frames": ("INT", {"default": 1, "min": 1, "max": 10000}) + } } CATEGORY = "image" @@ -208,12 +217,13 @@ class LoadImageMaskSequence: RETURN_TYPES = ("MASK_SEQUENCE",) FUNCTION = "load_image_sequence" - def load_image_sequence(self, image_sequence_folder, channel): + def load_image_sequence(self, image_sequence_folder, channel, sample_start_idx, sample_frame_rate, n_sample_frames): image_path = os.path.join(self.input_dir, image_sequence_folder) file_list = sorted(os.listdir(image_path), key=lambda s: sum(((s, int(n)) for s, n in re.findall(r'(\D+)(\d+)', 'a%s0' % s)), ())) sample_frames_mask = [] - for file in file_list: - i = Image.open(os.path.join(image_path, file)) + sample_index = list(range(sample_start_idx - 1, len(file_list), sample_frame_rate))[:n_sample_frames] + for num in sample_index: + i = Image.open(os.path.join(image_path, file_list[num])) mask = None c = channel[0].upper() if c in i.getbands(): diff --git a/tuneavideo/models/unet.py b/tuneavideo/models/unet.py index 20f02b0..4b05600 100644 --- a/tuneavideo/models/unet.py +++ b/tuneavideo/models/unet.py @@ -296,7 +296,6 @@ class UNet3DConditionModel(ModelMixin, ConfigMixin): context = context[0].unsqueeze(0) sample = rearrange(x.unsqueeze(0), "b f c h w -> b c f h w") - sample = sample.type(self.dtype) context = context.type(self.dtype) @@ -309,6 +308,9 @@ class UNet3DConditionModel(ModelMixin, ConfigMixin): down_block_additional_residuals.append(rearrange(output.unsqueeze(0), "a b c d e -> a c b d e")) mid_block_additional_residual = rearrange(control["middle"][0].unsqueeze(0), "a b c d e -> a c b d e") + del x, timesteps, control + torch.cuda.empty_cache() + # By default samples have to be AT least a multiple of the overall upsampling factor. # The overall upsampling factor is equal to 2 ** (# num of upsampling layears). # However, the upsampling interpolation output size can be forced to fit any upsampling size