diff --git a/lvdm/modules/attention.py b/lvdm/modules/attention.py index 15181f7..4a86de0 100644 --- a/lvdm/modules/attention.py +++ b/lvdm/modules/attention.py @@ -369,55 +369,77 @@ class TemporalTransformer(nn.Module): self.proj_out = zero_module(Linear(inner_dim, in_channels)) self.use_linear = use_linear - def forward(self, x, context=None): - b, c, t, h, w = x.shape - x_in = x - x = self.norm(x) - x = rearrange(x, 'b c t h w -> (b h w) c t').contiguous() - if not self.use_linear: - x = self.proj_in(x) - x = rearrange(x, 'bhw c t -> bhw t c').contiguous() - if self.use_linear: - x = self.proj_in(x) - temp_mask = None - if self.causal_attention: - # slice the from mask map - temp_mask = self.mask[:,:t,:t].to(x.device) + def forward(self, x_in, context=None, frame_window_size=None, frame_window_stride=None): + B, C, T, H, W = x_in.shape + def process_slice(x, t_start=None, t_end=None): + b, c, t, h, w = x.shape + x = self.norm(x) + x = rearrange(x, "b c t h w -> (b h w) c t").contiguous() + if not self.use_linear: + x = self.proj_in(x) + x = rearrange(x, "bhw c t -> bhw t c").contiguous() + if self.use_linear: + x = self.proj_in(x) - if temp_mask is not None: - mask = temp_mask.to(x.device) - mask = repeat(mask, 'l i j -> (l bhw) i j', bhw=b*h*w) + temp_mask = None + if self.causal_attention: + # slice the from mask map + if t_start is not None and t_end is not None: + temp_mask = self.mask[:, t_start:t_end, t_start:t_end].to(x.device) + else: + temp_mask = self.mask[:, :t, :t].to(x.device) + + if temp_mask is not None: + mask = temp_mask.to(x.device) + mask = repeat(mask, "l i j -> (l bhw) i j", bhw=b * h * w) + else: + mask = None + + if self.only_self_att: + ## note: if no context is given, cross-attention defaults to self-attention + for i, block in enumerate(self.transformer_blocks): + x = block(x, mask=mask) + x = rearrange(x, "(b hw) t c -> b hw t c", b=b).contiguous() + else: + x = rearrange(x, "(b hw) t c -> b hw t c", b=b).contiguous() + context = rearrange(context, "(b t) l con -> b t l con", t=t).contiguous() + for i, block in enumerate(self.transformer_blocks): + # calculate each batch one by one (since number in shape could not greater then 65,535 for some package) + for j in range(b): + context_j = repeat( + context[j], "t l con -> (t r) l con", r=(h * w) // t, t=t + ).contiguous() + ## note: causal mask will not applied in cross-attention case + x[j] = block(x[j], context=context_j) + + if self.use_linear: + x = self.proj_out(x) + x = rearrange(x, "b (h w) t c -> b c t h w", h=h, w=w).contiguous() + if not self.use_linear: + x = rearrange(x, "b hw t c -> (b hw) c t").contiguous() + x = self.proj_out(x) + x = rearrange(x, "(b h w) c t -> b c t h w", b=b, h=h, w=w).contiguous() + + return x + + if frame_window_size and frame_window_stride and T > frame_window_size: + views = get_frame_views(T, frame_window_size, frame_window_stride) + count = torch.zeros_like(x_in) + value = torch.zeros_like(x_in) + for t_start, t_end in views: + weight_sequence = get_frame_weight_sequence(t_end - t_start) + weight_tensor = torch.ones_like(count[:, :, t_start:t_end]) + weight_tensor = weight_tensor * torch.tensor(weight_sequence).to(x_in.device).unsqueeze(0).unsqueeze(-1).unsqueeze(-1) + x_slice = process_slice(x_in[:, :, t_start:t_end]) + value[:, :, t_start:t_end] += x_slice * weight_tensor + count[:, :, t_start:t_end] += weight_tensor + x = torch.where(count>0, value/count, value) else: - mask = None - - if self.only_self_att: - ## note: if no context is given, cross-attention defaults to self-attention - for i, block in enumerate(self.transformer_blocks): - x = block(x, mask=mask) - x = rearrange(x, '(b hw) t c -> b hw t c', b=b).contiguous() - else: - x = rearrange(x, '(b hw) t c -> b hw t c', b=b).contiguous() - context = rearrange(context, '(b t) l con -> b t l con', t=t).contiguous() - for i, block in enumerate(self.transformer_blocks): - # calculate each batch one by one (since number in shape could not greater then 65,535 for some package) - for j in range(b): - context_j = repeat( - context[j], - 't l con -> (t r) l con', r=(h * w) // t, t=t).contiguous() - ## note: causal mask will not applied in cross-attention case - x[j] = block(x[j], context=context_j) - - if self.use_linear: - x = self.proj_out(x) - x = rearrange(x, 'b (h w) t c -> b c t h w', h=h, w=w).contiguous() - if not self.use_linear: - x = rearrange(x, 'b hw t c -> (b hw) c t').contiguous() - x = self.proj_out(x) - x = rearrange(x, '(b h w) c t -> b c t h w', b=b, h=h, w=w).contiguous() + x = process_slice(x_in) return x + x_in - + class GEGLU(nn.Module): def __init__(self, dim_in, dim_out): @@ -519,3 +541,27 @@ class SpatialSelfAttention(nn.Module): h_ = self.proj_out(h_) return x+h_ + +def get_frame_views(video_length, window_size=16, stride=4): + """ + Gets frame views for context windowing + """ + num_blocks_time = (video_length - window_size) // stride + 1 + views = [] + for i in range(num_blocks_time): + t_start = int(i * stride) + t_end = t_start + window_size + views.append((t_start,t_end)) + return views + +def get_frame_weight_sequence(n): + """ + Gets a list of weights for merging context windows + """ + if n % 2 == 0: + max_weight = n // 2 + weight_sequence = list(range(1, max_weight + 1, 1)) + list(range(max_weight, 0, -1)) + else: + max_weight = (n + 1) // 2 + weight_sequence = list(range(1, max_weight, 1)) + [max_weight] + list(range(max_weight - 1, 0, -1)) + return weight_sequence diff --git a/lvdm/modules/networks/openaimodel3d.py b/lvdm/modules/networks/openaimodel3d.py index cb45a0f..a1b29b8 100644 --- a/lvdm/modules/networks/openaimodel3d.py +++ b/lvdm/modules/networks/openaimodel3d.py @@ -33,7 +33,7 @@ class TimestepEmbedSequential(nn.Sequential, TimestepBlock): support it as an extra input. """ - def forward(self, x, emb, context=None, batch_size=None): + def forward(self, x, emb, context=None, batch_size=None, frame_window_size=None, frame_window_stride=None): for layer in self: if isinstance(layer, TimestepBlock): x = layer(x, emb, batch_size=batch_size) @@ -41,7 +41,7 @@ class TimestepEmbedSequential(nn.Sequential, TimestepBlock): x = layer(x, context) elif isinstance(layer, TemporalTransformer): x = rearrange(x, '(b f) c h w -> b c f h w', b=batch_size) - x = layer(x, context) + x = layer(x, context, frame_window_size=frame_window_size, frame_window_stride=frame_window_stride) x = rearrange(x, 'b c f h w -> (b f) c h w') else: x = layer(x) @@ -545,7 +545,7 @@ class UNetModel(nn.Module): zero_module(conv_nd(dims, model_channels, out_channels, 3, padding=1)), ) - def forward(self, x, timesteps, context=None, features_adapter=None, fs=None, **kwargs): + def forward(self, x, timesteps, context=None, features_adapter=None, fs=None, frame_window_size=None, frame_window_stride=None, **kwargs): b,_,t,_,_ = x.shape t_emb = timestep_embedding(timesteps, self.model_channels, repeat_only=False).type(x.dtype) emb = self.time_embed(t_emb) @@ -580,9 +580,9 @@ class UNetModel(nn.Module): adapter_idx = 0 hs = [] for id, module in enumerate(self.input_blocks): - h = module(h, emb, context=context, batch_size=b) + h = module(h, emb, context=context, batch_size=b, frame_window_size=frame_window_size, frame_window_stride=frame_window_stride) if id ==0 and self.addition_attention: - h = self.init_attn(h, emb, context=context, batch_size=b) + h = self.init_attn(h, emb, context=context, batch_size=b, frame_window_size=frame_window_size, frame_window_stride=frame_window_stride) ## plug-in adapter features if ((id+1)%3 == 0) and features_adapter is not None: h = h + features_adapter[adapter_idx] @@ -591,13 +591,13 @@ class UNetModel(nn.Module): if features_adapter is not None: assert len(features_adapter)==adapter_idx, 'Wrong features_adapter' - h = self.middle_block(h, emb, context=context, batch_size=b) + h = self.middle_block(h, emb, context=context, batch_size=b, frame_window_size=frame_window_size, frame_window_stride=frame_window_stride) for module in self.output_blocks: h = torch.cat([h, hs.pop()], dim=1) - h = module(h, emb, context=context, batch_size=b) + h = module(h, emb, context=context, batch_size=b, frame_window_size=frame_window_size, frame_window_stride=frame_window_stride) h = h.type(x.dtype) y = self.out(h) # reshape back to (b c t h w) y = rearrange(y, '(b t) c h w -> b c t h w', b=b) - return y \ No newline at end of file + return y