Spaces:
Running on Zero
Running on Zero
File size: 11,515 Bytes
3e936b2 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 | import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.distributed as dist
from wan.modules.attention import attention
from wan.modules.model import WanRMSNorm
class MemoryCrossAttentionLayer(nn.Module):
def __init__(self, dim, num_heads, ffn_dim, qk_norm=True, eps=1e-06):
super().__init__()
self.num_heads = num_heads
self.head_dim = dim // num_heads
self.q = nn.Linear(dim, dim)
self.o = nn.Linear(dim, dim)
self.norm_q = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
self.norm1 = WanRMSNorm(dim, eps=eps)
self.norm2 = WanRMSNorm(dim, eps=eps)
self.ffn = nn.Sequential(nn.Linear(dim, ffn_dim), nn.GELU(approximate='tanh'), nn.Linear(ffn_dim, dim))
def forward(self, query_state, cached_k, cached_v):
b = query_state.shape[0]
m = query_state.shape[1]
n, d = (self.num_heads, self.head_dim)
residual = query_state
h = self.norm1(query_state)
q = self.norm_q(self.q(h)).view(b, m, n, d)
out = attention(q, cached_k, cached_v)
out = out.flatten(2).type_as(h)
out = self.o(out)
query_state = residual + out
residual = query_state
query_state = residual + self.ffn(self.norm2(query_state))
return query_state
class QueryMemoryEncoder(nn.Module):
def __init__(self, config):
super().__init__()
Q_frames = getattr(config, 'Q_frames', 3)
tokens_per_frame = getattr(config, 'tokens_per_frame', 1560)
n_encoder_layers = getattr(config, 'n_encoder_layers', 2)
hidden_dim = getattr(config, 'hidden_dim', 1536)
num_heads = getattr(config, 'num_heads', 12)
head_dim = getattr(config, 'head_dim', 128)
ffn_dim = hidden_dim * 4
gate_init_bias = getattr(config, 'gate_init_bias', 2.0)
qk_norm = getattr(config, 'qk_norm', True)
eps = 1e-06
self.M = Q_frames * tokens_per_frame
self.hidden_dim = hidden_dim
self.num_heads = num_heads
self.head_dim = head_dim
self.Q_frames = Q_frames
self.tokens_per_frame = tokens_per_frame
self.use_batch_update = getattr(config, 'use_batch_update', False)
self.batch_update_interval = getattr(config, 'batch_update_interval', 1)
self.use_sink_anchor = getattr(config, 'use_sink_anchor', False)
self.use_vib = getattr(config, 'use_vib', False)
self.bptt_clips = getattr(config, 'bptt_clips', 1)
self.encoder_lr_multiplier = getattr(config, 'encoder_lr_multiplier', 5.0)
self.normalize_memory_k = getattr(config, 'normalize_memory_k', False)
self.use_residual_update = getattr(config, 'use_residual_update', False)
self.use_post_norm = getattr(config, 'use_post_norm', False)
self.memory_recache = getattr(config, 'memory_recache', False)
self.num_query_groups = getattr(config, 'num_query_groups', 1)
initializer_range = getattr(config, 'initializer_range', 0.014)
self.layers = nn.ModuleList([MemoryCrossAttentionLayer(hidden_dim, num_heads, ffn_dim, qk_norm, eps) for _ in range(n_encoder_layers)])
if self.use_post_norm:
self.post_norm = WanRMSNorm(hidden_dim, eps=eps)
if self.num_query_groups > 1:
self.query_inits = nn.ParameterList([nn.Parameter(torch.randn(1, self.M, hidden_dim) * initializer_range) for _ in range(self.num_query_groups)])
self.to_k_groups = nn.ModuleList([nn.Linear(hidden_dim, num_heads * head_dim) for _ in range(self.num_query_groups)])
self.to_v_groups = nn.ModuleList([nn.Linear(hidden_dim, num_heads * head_dim) for _ in range(self.num_query_groups)])
if self.normalize_memory_k:
self.norm_k_out_groups = nn.ModuleList([WanRMSNorm(num_heads * head_dim, eps=eps) for _ in range(self.num_query_groups)])
self.connector_projs = nn.ModuleList([nn.Sequential(nn.Linear(hidden_dim, hidden_dim), nn.GELU(approximate='tanh'), nn.Linear(hidden_dim, hidden_dim), WanRMSNorm(hidden_dim, eps=eps)) for _ in range(self.num_query_groups)])
self.gate_linears = nn.ModuleList([nn.Linear(hidden_dim * 2, hidden_dim) for _ in range(self.num_query_groups)])
for gl in self.gate_linears:
with torch.no_grad():
gl.bias.fill_(gate_init_bias)
else:
self.query_init = nn.Parameter(torch.randn(1, self.M, hidden_dim) * initializer_range)
self.connector_proj = nn.Sequential(nn.Linear(hidden_dim, hidden_dim), nn.GELU(approximate='tanh'), nn.Linear(hidden_dim, hidden_dim), WanRMSNorm(hidden_dim, eps=eps))
self.gate_linear = nn.Linear(hidden_dim * 2, hidden_dim)
with torch.no_grad():
self.gate_linear.bias.fill_(gate_init_bias)
self.to_k = nn.Linear(hidden_dim, num_heads * head_dim)
self.to_v = nn.Linear(hidden_dim, num_heads * head_dim)
if self.normalize_memory_k:
self.norm_k_out = WanRMSNorm(num_heads * head_dim, eps=eps)
if self.use_vib:
self.mu_proj = nn.Linear(hidden_dim, hidden_dim)
self.logvar_proj = nn.Linear(hidden_dim, hidden_dim)
self._evicted_k_buffer = []
self._evicted_v_buffer = []
self.query_state = None
self.has_history = False
self._update_count = 0
total_params = sum((p.numel() for p in self.parameters()))
def reset(self, batch_size=1, device=None, dtype=None):
def _expand_init(init):
if init.dim() != 3 or init.shape[1] != self.M:
raise RuntimeError(f'query_init has unexpected shape {list(init.shape)} (expected [1, {self.M}, {self.hidden_dim}]). This usually means the encoder was placed inside FSDP, which flattens parameters. The encoder must remain outside FSDP (attached via object.__setattr__).')
state = init.expand(batch_size, -1, -1).clone()
if device is not None:
state = state.to(device=device)
if dtype is not None:
state = state.to(dtype=dtype)
return state
if self.num_query_groups > 1:
self.query_states = [_expand_init(init) for init in self.query_inits]
self.query_state = self.query_states[0]
else:
self.query_state = _expand_init(self.query_init)
self.has_history = False
self._update_count = 0
self._evicted_k_buffer = []
self._evicted_v_buffer = []
shape = list(self.query_states[0].shape) if self.num_query_groups > 1 else list(self.query_state.shape)
def update(self, evicted_k, evicted_v, sink_k=None, sink_v=None):
if self.query_state is None:
self.reset(batch_size=evicted_k.shape[0], device=evicted_k.device, dtype=evicted_k.dtype if evicted_k.dtype.is_floating_point else torch.bfloat16)
query_state = self.query_state
if self.use_batch_update:
self._evicted_k_buffer.append(evicted_k)
self._evicted_v_buffer.append(evicted_v)
if len(self._evicted_k_buffer) < self.batch_update_interval:
return 0.0
evicted_k = torch.cat(self._evicted_k_buffer, dim=1)
evicted_v = torch.cat(self._evicted_v_buffer, dim=1)
self._evicted_k_buffer = []
self._evicted_v_buffer = []
if self.use_sink_anchor and sink_k is not None:
ctx_k = torch.cat([sink_k, evicted_k], dim=1)
ctx_v = torch.cat([sink_v, evicted_v], dim=1)
else:
ctx_k = evicted_k
ctx_v = evicted_v
def _update_single_group(qs, connector_proj, gate_linear):
state = qs
for layer in self.layers:
state = layer(state, ctx_k, ctx_v)
if self.use_residual_update:
new_state = state
gate_mean = None
else:
projected = connector_proj(state)
gate = torch.sigmoid(gate_linear(torch.cat([qs, projected], dim=-1)))
new_state = gate * qs + (1 - gate) * projected
gate_mean = gate.mean().item()
if self.use_post_norm:
new_state = self.post_norm(new_state)
return (new_state, gate_mean)
kl_loss = 0.0
if self.num_query_groups > 1:
gate_means = []
for g in range(self.num_query_groups):
new_state, gm = _update_single_group(self.query_states[g], self.connector_projs[g], self.gate_linears[g])
if self.use_vib and self.training:
mu = self.mu_proj(new_state)
logvar = self.logvar_proj(new_state)
std = torch.exp(0.5 * logvar)
new_state = mu + torch.randn_like(std) * std
kl_loss += -0.5 * (1 + logvar - mu.pow(2) - logvar.exp()).sum()
elif self.use_vib:
new_state = self.mu_proj(new_state)
self.query_states[g] = new_state
gate_means.append(gm)
self.query_state = self.query_states[0]
gate_str = ', '.join((f"g{i}={('res' if gm is None else f'{gm:.3f}')}" for i, gm in enumerate(gate_means)))
norm_str = ', '.join((f'g{i}={self.query_states[i].norm().item():.1f}' for i in range(self.num_query_groups)))
else:
new_state, gate_mean = _update_single_group(query_state, self.connector_proj, self.gate_linear)
if self.use_vib and self.training:
mu = self.mu_proj(new_state)
logvar = self.logvar_proj(new_state)
std = torch.exp(0.5 * logvar)
new_state = mu + torch.randn_like(std) * std
kl_loss = -0.5 * (1 + logvar - mu.pow(2) - logvar.exp()).sum()
elif self.use_vib:
new_state = self.mu_proj(new_state)
self.query_state = new_state
gate_str = f"{('residual_update' if gate_mean is None else f'gate_mean={gate_mean:.3f}')}"
norm_str = f'{new_state.norm().item():.1f}'
self.has_history = True
self._update_count += 1
return kl_loss
def get_kv(self, group_index=None):
if not self.has_history:
return None
if self.num_query_groups > 1 and group_index is not None:
state = self.query_states[group_index]
k = self.to_k_groups[group_index](state)
if self.normalize_memory_k:
k = self.norm_k_out_groups[group_index](k)
k = k.view(state.shape[0], self.M, self.num_heads, self.head_dim)
v = self.to_v_groups[group_index](state).view(state.shape[0], self.M, self.num_heads, self.head_dim)
return (k, v)
else:
state = self.query_state
k = self.to_k(state)
if self.normalize_memory_k:
k = self.norm_k_out(k)
k = k.view(state.shape[0], self.M, self.num_heads, self.head_dim)
v = self.to_v(state).view(state.shape[0], self.M, self.num_heads, self.head_dim)
return (k, v)
def detach_state(self):
if self.num_query_groups > 1:
self.query_states = [s.detach() for s in self.query_states]
self.query_state = self.query_states[0]
elif self.query_state is not None:
self.query_state = self.query_state.detach()
|