add context windows

This commit is contained in:
Benjamin Paine
2024-03-19 11:13:34 -04:00
parent c6178d9fdc
commit 53c2292f16
2 changed files with 98 additions and 52 deletions
+90 -44
View File
@@ -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
+8 -8
View File
@@ -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
return y