first modification

This commit is contained in:
sylym
2023-03-24 21:46:38 +08:00
parent ebd653b385
commit 9c2998dc2b
2 changed files with 21 additions and 9 deletions
+18 -8
View File
@@ -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():
+3 -1
View File
@@ -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