cleanup
This commit is contained in:
+1
-146
@@ -1677,7 +1677,7 @@ class WanModel(torch.nn.Module):
|
||||
attention_mode=attention_mode
|
||||
)
|
||||
self.trainable_cond_mask = nn.Embedding(3, self.dim)
|
||||
|
||||
|
||||
self.frame_packer = FramePackMotioner(
|
||||
inner_dim=self.dim,
|
||||
num_heads=self.num_heads,
|
||||
@@ -1838,151 +1838,6 @@ class WanModel(torch.nn.Module):
|
||||
block.to(self.offload_device, non_blocking=self.use_non_blocking)
|
||||
|
||||
return hints
|
||||
|
||||
def process_motion(self, motion_latents, drop_motion_frames=False):
|
||||
if drop_motion_frames or motion_latents[0].shape[1] == 0:
|
||||
return [], []
|
||||
self.lat_motion_frames = motion_latents[0].shape[1]
|
||||
mot = [self.patch_embedding(m.unsqueeze(0)) for m in motion_latents]
|
||||
batch_size = len(mot)
|
||||
|
||||
mot_remb = []
|
||||
flattern_mot = []
|
||||
for bs in range(batch_size):
|
||||
height, width = mot[bs].shape[3], mot[bs].shape[4]
|
||||
flat_mot = mot[bs].flatten(2).transpose(1, 2).contiguous()
|
||||
motion_grid_sizes = [[
|
||||
torch.tensor([-self.lat_motion_frames, 0,
|
||||
0]).unsqueeze(0).repeat(1, 1),
|
||||
torch.tensor([0, height, width]).unsqueeze(0).repeat(1, 1),
|
||||
torch.tensor([self.lat_motion_frames, height,
|
||||
width]).unsqueeze(0).repeat(1, 1)
|
||||
]]
|
||||
motion_rope_emb = rope_precompute(
|
||||
flat_mot.detach().view(1, flat_mot.shape[1], self.num_heads,
|
||||
self.dim // self.num_heads),
|
||||
motion_grid_sizes,
|
||||
self.freqs,
|
||||
start=None)
|
||||
mot_remb.append(motion_rope_emb)
|
||||
flattern_mot.append(flat_mot)
|
||||
return flattern_mot, mot_remb
|
||||
|
||||
def process_motion_frame_pack(self,
|
||||
motion_latents,
|
||||
drop_motion_frames=False,
|
||||
add_last_motion=2):
|
||||
flattern_mot, mot_remb = self.frame_packer(motion_latents,
|
||||
add_last_motion)
|
||||
if drop_motion_frames:
|
||||
return [m[:, :0] for m in flattern_mot
|
||||
], [m[:, :0] for m in mot_remb]
|
||||
else:
|
||||
return flattern_mot, mot_remb
|
||||
|
||||
def process_motion_transformer_motioner(self,
|
||||
motion_latents,
|
||||
drop_motion_frames=False,
|
||||
add_last_motion=True):
|
||||
batch_size, height, width = len(
|
||||
motion_latents), motion_latents[0].shape[2] // self.patch_size[
|
||||
1], motion_latents[0].shape[3] // self.patch_size[2]
|
||||
|
||||
freqs = self.freqs
|
||||
device = self.patch_embedding.weight.device
|
||||
if freqs.device != device:
|
||||
freqs = freqs.to(device)
|
||||
if self.trainable_token_pos_emb:
|
||||
token_freqs = self.token_freqs.to(torch.float64)
|
||||
token_freqs = token_freqs / token_freqs.norm(
|
||||
dim=-1, keepdim=True)
|
||||
freqs = [freqs, torch.view_as_complex(token_freqs)]
|
||||
|
||||
if not drop_motion_frames and add_last_motion:
|
||||
last_motion_latent = [u[:, -1:] for u in motion_latents]
|
||||
last_mot = [
|
||||
self.patch_embedding(m.unsqueeze(0)) for m in last_motion_latent
|
||||
]
|
||||
last_mot = [m.flatten(2).transpose(1, 2) for m in last_mot]
|
||||
last_mot = torch.cat(last_mot)
|
||||
gride_sizes = [[
|
||||
torch.tensor([-1, 0, 0]).unsqueeze(0).repeat(batch_size, 1),
|
||||
torch.tensor([0, height,
|
||||
width]).unsqueeze(0).repeat(batch_size, 1),
|
||||
torch.tensor([1, height,
|
||||
width]).unsqueeze(0).repeat(batch_size, 1)
|
||||
]]
|
||||
else:
|
||||
last_mot = torch.zeros([batch_size, 0, self.dim],
|
||||
device=motion_latents[0].device,
|
||||
dtype=motion_latents[0].dtype)
|
||||
gride_sizes = []
|
||||
|
||||
zip_motion = self.motioner(motion_latents)
|
||||
zip_motion = self.zip_motion_out(zip_motion)
|
||||
if drop_motion_frames:
|
||||
zip_motion = zip_motion * 0.0
|
||||
zip_motion_grid_sizes = [[
|
||||
torch.tensor([-1, 0, 0]).unsqueeze(0).repeat(batch_size, 1),
|
||||
torch.tensor([
|
||||
0, self.motioner.motion_side_len, self.motioner.motion_side_len
|
||||
]).unsqueeze(0).repeat(batch_size, 1),
|
||||
torch.tensor(
|
||||
[1 if not self.trainable_token_pos_emb else -1, height,
|
||||
width]).unsqueeze(0).repeat(batch_size, 1),
|
||||
]]
|
||||
|
||||
mot = torch.cat([last_mot, zip_motion], dim=1)
|
||||
gride_sizes = gride_sizes + zip_motion_grid_sizes
|
||||
|
||||
motion_rope_emb = rope_precompute(
|
||||
mot.detach().view(batch_size, mot.shape[1], self.num_heads,
|
||||
self.dim // self.num_heads),
|
||||
gride_sizes,
|
||||
freqs,
|
||||
start=None)
|
||||
return [m.unsqueeze(0) for m in mot
|
||||
], [r.unsqueeze(0) for r in motion_rope_emb]
|
||||
|
||||
def inject_motion(self,
|
||||
x,
|
||||
seq_lens,
|
||||
rope_embs,
|
||||
mask_input,
|
||||
motion_latents,
|
||||
drop_motion_frames=False,
|
||||
add_last_motion=True):
|
||||
# inject the motion frames token to the hidden states
|
||||
if self.enable_motioner:
|
||||
mot, mot_remb = self.process_motion_transformer_motioner(
|
||||
motion_latents,
|
||||
drop_motion_frames=drop_motion_frames,
|
||||
add_last_motion=add_last_motion)
|
||||
elif self.enable_framepack:
|
||||
mot, mot_remb = self.process_motion_frame_pack(
|
||||
motion_latents,
|
||||
drop_motion_frames=drop_motion_frames,
|
||||
add_last_motion=add_last_motion)
|
||||
else:
|
||||
mot, mot_remb = self.process_motion(
|
||||
motion_latents, drop_motion_frames=drop_motion_frames)
|
||||
|
||||
if len(mot) > 0:
|
||||
x = [torch.cat([u, m], dim=1) for u, m in zip(x, mot)]
|
||||
seq_lens = seq_lens + torch.tensor([r.size(1) for r in mot],
|
||||
dtype=torch.long)
|
||||
rope_embs = [
|
||||
torch.cat([u, m], dim=1) for u, m in zip(rope_embs, mot_remb)
|
||||
]
|
||||
mask_input = [
|
||||
torch.cat([
|
||||
m, 2 * torch.ones([1, u.shape[1] - m.shape[1]],
|
||||
device=m.device,
|
||||
dtype=m.dtype)
|
||||
],
|
||||
dim=1) for m, u in zip(mask_input, x)
|
||||
]
|
||||
return x, seq_lens, rope_embs, mask_input
|
||||
|
||||
def audio_injector_forward(self, block_idx, x, audio_emb, scale=1.0):
|
||||
if block_idx in self.audio_injector.injected_block_id.keys():
|
||||
|
||||
Reference in New Issue
Block a user