rope tweaks

This commit is contained in:
kijai
2025-03-13 02:26:40 +02:00
parent 36a0b0fb25
commit 71fd6c7208
2 changed files with 47 additions and 31 deletions
+17 -3
View File
@@ -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
View File
@@ -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