rope tweaks
This commit is contained in:
@@ -1097,6 +1097,7 @@ class WanVideoContextOptions:
|
||||
"optional": {
|
||||
"image_cond_start_step": ("INT", {"default": 6, "min": 0, "max": 10000, "step": 1, "tooltip": "!EXPERIMENTAL! Start step of using previous window results as input instead of the init image"}),
|
||||
"image_cond_window_count": ("INT", {"default": 2, "min": 1, "max": 10000, "step": 1, "tooltip": "!EXPERIMENTAL! Number of image 'prompt windows'"}),
|
||||
"vae": ("WANVAE",),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1106,7 +1107,7 @@ class WanVideoContextOptions:
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
DESCRIPTION = "Context options for WanVideo, allows splitting the video into context windows and attemps blending them for longer generations than the model and memory otherwise would allow."
|
||||
|
||||
def process(self, context_schedule, context_frames, context_stride, context_overlap, freenoise, verbose, image_cond_start_step=6, image_cond_window_count=2):
|
||||
def process(self, context_schedule, context_frames, context_stride, context_overlap, freenoise, verbose, image_cond_start_step=6, image_cond_window_count=2, vae=None):
|
||||
context_options = {
|
||||
"context_schedule":context_schedule,
|
||||
"context_frames":context_frames,
|
||||
@@ -1115,7 +1116,8 @@ class WanVideoContextOptions:
|
||||
"freenoise":freenoise,
|
||||
"verbose":verbose,
|
||||
"image_cond_start_step": image_cond_start_step,
|
||||
"image_cond_window_count": image_cond_window_count
|
||||
"image_cond_window_count": image_cond_window_count,
|
||||
"vae": vae,
|
||||
}
|
||||
|
||||
return (context_options,)
|
||||
@@ -1307,6 +1309,9 @@ class WanVideoSampler:
|
||||
context_frames = (context_options["context_frames"] - 1) // 4 + 1
|
||||
context_stride = context_options["context_stride"] // 4
|
||||
context_overlap = context_options["context_overlap"] // 4
|
||||
context_vae = context_options.get("vae", None)
|
||||
if context_vae is not None:
|
||||
context_vae.to(device)
|
||||
|
||||
self.window_tracker = WindowTracker(verbose=context_options["verbose"])
|
||||
|
||||
@@ -1725,11 +1730,20 @@ class WanVideoSampler:
|
||||
if idx >= context_options["image_cond_start_step"]:
|
||||
#strength = 0.5
|
||||
#partial_image_cond *= strength
|
||||
if context_vae is not None:
|
||||
to_decode = self.previous_noise_pred_context[:,-1,:, :].unsqueeze(1).unsqueeze(0).to(context_vae.dtype)
|
||||
#to_decode = to_decode.permute(0, 1, 3, 2)
|
||||
print("to_decode.shape", to_decode.shape)
|
||||
image = context_vae.decode(to_decode, device=device, tiled=False)[0]
|
||||
print("decoded image.shape", image.shape) #torch.Size([3, 37, 832, 480])
|
||||
image = context_vae.encode(image.unsqueeze(0).to(context_vae.dtype), device=device, tiled=False)
|
||||
print("encoded image.shape", image.shape)
|
||||
#partial_img_emb[:, 0, :, :] = image[0][:,0,:,:]
|
||||
print("partial_img_emb.shape", partial_img_emb.shape)
|
||||
mask = torch.ones(4, partial_img_emb.shape[2], partial_img_emb.shape[3], device=partial_img_emb.device, dtype=partial_img_emb.dtype) #torch.Size([20, 10, 104, 60])
|
||||
print("mask.shape", mask.shape)
|
||||
print("self.previous_noise_pred_context.shape", self.previous_noise_pred_context.shape) #torch.Size([16, 10, 104, 60])
|
||||
partial_img_emb[:, 0, :, :] = torch.cat([self.previous_noise_pred_context[:, -1, :, :], mask], dim=0)
|
||||
partial_img_emb[:, 0, :, :] = torch.cat([image[0][:,0,:,:], mask], dim=0)
|
||||
else:
|
||||
partial_img_emb[:, 0, :, :] = partial_image_cond
|
||||
|
||||
|
||||
+30
-28
@@ -52,44 +52,46 @@ def rope_params(max_seq_len, dim, theta=10000, L_test=25, k=0):
|
||||
|
||||
from comfy.model_management import get_torch_device, get_autocast_device
|
||||
@torch.autocast(device_type=get_autocast_device(get_torch_device()), enabled=False)
|
||||
@torch.compiler.disable()
|
||||
def rope_apply(x, grid_sizes, freqs):
|
||||
batch_size, max_seq, n, c_total = x.shape
|
||||
c = c_total // 2
|
||||
n, c = x.size(2), x.size(3) // 2
|
||||
|
||||
# Static splits
|
||||
c1 = c - 2 * (c // 3)
|
||||
c2 = c // 3
|
||||
c3 = c // 3
|
||||
freqs_split = freqs.split([c1, c2, c3], dim=1)
|
||||
|
||||
def process_chunk(x_chunk, f, h, w, freqs_split):
|
||||
seq_len = f * h * w
|
||||
x_complex = torch.view_as_complex(x_chunk.to(torch.float64).reshape(seq_len, n, c, 2))
|
||||
|
||||
f1 = freqs_split[0][:f].view(f, 1, 1, -1)
|
||||
f2 = freqs_split[1][:h].view(1, h, 1, -1)
|
||||
f3 = freqs_split[2][:w].view(1, 1, w, -1)
|
||||
|
||||
freq_cat = torch.cat([
|
||||
f1.repeat(1, h, w, 1),
|
||||
f2.repeat(f, 1, w, 1),
|
||||
f3.repeat(f, h, 1, 1)
|
||||
], dim=-1).reshape(seq_len, 1, c)
|
||||
|
||||
x_i = torch.view_as_real(x_complex * freq_cat)
|
||||
return x_i.reshape(seq_len, n, c_total)
|
||||
freqs = freqs.split([c - 2 * (c // 3), c // 3, c // 3], dim=1)
|
||||
|
||||
output = []
|
||||
for i, (f, h, w) in enumerate(grid_sizes.tolist()):
|
||||
seq_len = f * h * w
|
||||
x_i = process_chunk(x[i, :seq_len], f, h, w, freqs_split)
|
||||
|
||||
if seq_len < max_seq:
|
||||
@torch.compiler.disable()
|
||||
def view_as_complex_no_compile(x):
|
||||
x_i = torch.view_as_complex(x[i, :seq_len].to(torch.float64).reshape(seq_len, n, -1, 2))
|
||||
return x_i
|
||||
|
||||
x_i = view_as_complex_no_compile(x)
|
||||
|
||||
f_size = (f, 1, 1, -1)
|
||||
h_size = (1, h, 1, -1)
|
||||
w_size = (1, 1, w, -1)
|
||||
|
||||
freq_cat = torch.cat([
|
||||
freqs[0][:f].view(*f_size).expand(f, h, w, -1),
|
||||
freqs[1][:h].view(*h_size).expand(f, h, w, -1),
|
||||
freqs[2][:w].view(*w_size).expand(f, h, w, -1)
|
||||
], dim=-1).reshape(seq_len, 1, -1)
|
||||
|
||||
@torch.compiler.disable()
|
||||
def view_as_real_no_compile(x_i):
|
||||
x_i.mul_(freq_cat)
|
||||
x_i = torch.view_as_real(x_i).flatten(2)
|
||||
return x_i
|
||||
|
||||
x_i = view_as_real_no_compile(x_i)
|
||||
del freq_cat
|
||||
|
||||
if seq_len < x.size(1):
|
||||
x_i = torch.cat([x_i, x[i, seq_len:]], dim=0)
|
||||
output.append(x_i)
|
||||
|
||||
return torch.stack(output)
|
||||
return torch.stack(output).to(torch.float32)
|
||||
|
||||
def rope_apply_original(x, grid_sizes, freqs):
|
||||
n, c = x.size(2), x.size(3) // 2
|
||||
|
||||
Reference in New Issue
Block a user