diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index bfc8cb6..4813b54 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -65,7 +65,7 @@ def rope_apply(x, grid_sizes, freqs): def process_chunk(x_chunk, f, h, w, freqs_split): seq_len = f * h * w - x_complex = torch.view_as_complex(x_chunk.reshape(seq_len, n, c, 2)) + 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) @@ -91,6 +91,35 @@ def rope_apply(x, grid_sizes, freqs): return torch.stack(output) +def rope_apply_original(x, grid_sizes, freqs): + n, c = x.size(2), x.size(3) // 2 + + # split freqs + freqs = freqs.split([c - 2 * (c // 3), c // 3, c // 3], dim=1) + + # loop over samples + output = [] + for i, (f, h, w) in enumerate(grid_sizes.tolist()): + seq_len = f * h * w + + # precompute multipliers + x_i = torch.view_as_complex(x[i, :seq_len].to(torch.float64).reshape( + seq_len, n, -1, 2)) + freqs_i = torch.cat([ + freqs[0][:f].view(f, 1, 1, -1).expand(f, h, w, -1), + freqs[1][:h].view(1, h, 1, -1).expand(f, h, w, -1), + freqs[2][:w].view(1, 1, w, -1).expand(f, h, w, -1) + ], + dim=-1).reshape(seq_len, 1, -1) + + # apply rotary embedding + x_i = torch.view_as_real(x_i * freqs_i).flatten(2) + x_i = torch.cat([x_i, x[i, seq_len:]]) + + # append to collection + output.append(x_i) + return torch.stack(output).float() + class WanRMSNorm(nn.Module):