aalfekka commited on
Commit
c539dd8
·
1 Parent(s): 2e9ae79

added source code

Browse files
Files changed (2) hide show
  1. language_modeling.py +476 -0
  2. language_modeling_amt.py +1157 -0
language_modeling.py ADDED
@@ -0,0 +1,476 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # copied from original ARTM repo: https://raw.githubusercontent.com/RodkinIvan/associative-recurrent-memory-transformer/refs/heads/framework_accel/modeling_amt/language_modeling.py
2
+ import math
3
+ import torch
4
+ from torch.nn import CrossEntropyLoss
5
+ from transformers.modeling_outputs import CausalLMOutputWithCrossAttentions
6
+ from torch.nn.functional import relu as r
7
+
8
+ def dpfp(x, nu=1):
9
+ x = torch.cat([r(x), r(-x)], dim=-1)
10
+ x_rolled = torch.cat([x.roll(shifts=j, dims=-1)
11
+ for j in range(1,nu+1)], dim=-1)
12
+ x_repeat = torch.cat([x] * nu, dim=-1)
13
+ return x_repeat * x_rolled
14
+
15
+ class DPFP:
16
+ def __init__(self, nu):
17
+ self.nu = nu
18
+
19
+ def __call__(self, x):
20
+ nu = self.nu
21
+ x = torch.cat([r(x), r(-x)], dim=-1)
22
+ x_rolled = torch.cat([x.roll(shifts=j, dims=-1) for j in range(1,nu+1)], dim=-1)
23
+ x_repeat = torch.cat([x] * nu, dim=-1)
24
+ return x_repeat * x_rolled
25
+
26
+ class AssociativeLayerWrapper(torch.nn.Module):
27
+
28
+ def __init__(self, layer, d_model, num_mem_tokens, d_mem, correction=True, info=None) -> None:
29
+ super().__init__()
30
+ self.info = info
31
+ self.seg_num = 0
32
+ self.d_model = d_model
33
+ self.num_mem_tokens = num_mem_tokens
34
+ self.d_mem = d_mem
35
+
36
+ nu = 3
37
+ self.d_key = 2 * nu * d_mem
38
+ self.phi = DPFP(nu)
39
+ # self.d_key = d_mem
40
+ # self.phi = torch.nn.Identity()
41
+
42
+ self.W_mq = torch.nn.Linear(d_model, d_mem, bias=False, dtype=torch.bfloat16)
43
+ # torch.nn.init.zeros_(self.W_mq.weight)
44
+ self.W_mk = torch.nn.Linear(d_model, d_mem, bias=False, dtype=torch.bfloat16)
45
+ self.W_mv = torch.nn.Linear(d_model, d_model, bias=False, dtype=torch.bfloat16)
46
+ torch.nn.init.zeros_(self.W_mv.weight)
47
+ self.W_mb = torch.nn.Linear(d_model, 1, dtype=torch.bfloat16)
48
+
49
+ self.W_mem = torch.zeros(1, self.d_key, d_model, dtype=torch.bfloat16)
50
+ self.z = torch.zeros(1, self.d_key, dtype=torch.bfloat16)
51
+ self.W_mem.requires_grad_(False)
52
+ self.z.requires_grad_(False)
53
+
54
+ # self.ln = torch.nn.LayerNorm(d_model)
55
+
56
+ self.zero_mem()
57
+
58
+ self.layer = layer
59
+
60
+ self.generate_mode = False
61
+ self.first_seg = True
62
+ self.correction = correction
63
+
64
+ def associate(self, hidden_states):
65
+
66
+ self.W_mem = self.W_mem.to(hidden_states.device).to(torch.bfloat16)
67
+ self.z = self.z.to(hidden_states.device).to(torch.bfloat16)
68
+
69
+ mq = self.phi(self.W_mq(hidden_states)).to(torch.bfloat16) # (bsz, seq_len, 2d_mem * nu)
70
+
71
+ # crutch for dataparallel
72
+ # mq += 0 * self.W_mb(hidden_states).sum() * self.W_mk(hidden_states).sum() * self.W_mv(hidden_states).sum()
73
+ #print(mq, self.W_mem)
74
+ #print(mq.dtype, self.W_mem.dtype)
75
+ num = torch.einsum('ijk,ikt->ijt', mq, self.W_mem)
76
+ denom = torch.einsum("ik,ijk->ij", self.z, mq)[..., None] + 1e-5
77
+ hidden_states = num / denom
78
+
79
+ return hidden_states
80
+
81
+ def forward(self, hidden_states, **kwargs):
82
+ if not self.first_seg:
83
+ hidden_states = self.associate(
84
+ # self.ln(
85
+ hidden_states
86
+ # )
87
+ ) + hidden_states
88
+ out = self.layer(hidden_states=hidden_states, **kwargs)
89
+ if not self.generate_mode:
90
+ mem_tokens = out[0][:, -self.num_mem_tokens:]
91
+ self.update_mem(mem_tokens)
92
+ self.first_seg = False
93
+ return out
94
+
95
+ def update_mem(self, mem_tokens):
96
+
97
+ self.W_mem = self.W_mem.to(mem_tokens.device)
98
+ self.z = self.z.to(mem_tokens.device)
99
+
100
+ mk = self.phi(self.W_mk(mem_tokens))
101
+ new_mv = self.W_mv(mem_tokens) # (bsz, num_mem_tokens, d_model)
102
+ if not self.first_seg:
103
+ num = torch.einsum('ijk,ikt->ijt', mk, self.W_mem)
104
+ denom = torch.einsum("ij,ikj->ik", self.z, mk)[..., None] + 1e-5
105
+ prev_mv = num / denom
106
+ if self.correction:
107
+ new_info_coef = 1 - denom / (torch.linalg.norm(mk, dim=-1) ** 2 + 1e-5)[..., None]
108
+ new_info_coef = torch.clip(new_info_coef, 0, 1).detach()
109
+ else:
110
+ new_info_coef = 1
111
+ else:
112
+ prev_mv = torch.zeros_like(new_mv, device=new_mv.device)
113
+ new_info_coef = 1
114
+
115
+ # wandb.log({f"gamma_{self.info['layer']}": new_info_coef.mean(dim=1).item() if isinstance(new_info_coef, torch.Tensor) else 1}, step=self.seg_num)
116
+ mv = new_mv - prev_mv
117
+
118
+ # new_norm = torch.linalg.norm(new_mv, dim=-1)
119
+ # old_norm = torch.linalg.norm(prev_mv, dim=-1)
120
+ # new_info_coef = torch.clip(1 - old_norm / (new_norm + 1e-5), -10, 10)[..., None].detach()
121
+ # new_info_coef = 1 - denom
122
+
123
+ mb = torch.sigmoid(self.W_mb(mem_tokens))[..., 0]
124
+
125
+ associations = torch.einsum('ijk,ijt,ij->ikt', mk, mv, mb) # (bsz, d_mem, d_model)
126
+ self.W_mem = self.W_mem + associations
127
+
128
+ self.z = self.z + (new_info_coef*mk).sum(dim=1)
129
+ # self.z = self.z + (new_info_coef*mb[..., None]*mk).sum(dim=1)
130
+ self.seg_num += 1
131
+
132
+
133
+ def zero_mem(self):
134
+ self.first_seg = True
135
+ self.W_mem = torch.zeros(1, self.d_key, self.d_model)
136
+ self.z = torch.zeros(1, self.d_key)
137
+ self.seg_num = 0
138
+
139
+
140
+
141
+ class AssociativeMemoryCell(torch.nn.Module):
142
+ def __init__(self, base_model, num_mem_tokens, d_mem, layers_attr: str = 'transformer.h', wrap_pos=True, correction=True, use_lora=False, attend_to_previous_input=False):
143
+ super().__init__()
144
+ self.model = base_model
145
+ self.attend_to_previous_input = attend_to_previous_input
146
+ self.previous_input = None
147
+ self.num_mem_tokens = num_mem_tokens
148
+ self.d_mem = d_mem
149
+ self.d_model = base_model.get_input_embeddings().embedding_dim
150
+ self.W_mq = torch.nn.ModuleList()
151
+ self.W_mem = []
152
+ if use_lora:
153
+ # LoRA case
154
+ self.layers = self.model.model
155
+ else:
156
+ self.layers = self.model
157
+
158
+ self.layers_attrs = layers_attr.split('.')
159
+ for i, attr in enumerate(self.layers_attrs):
160
+ self.layers = getattr(self.layers, attr)
161
+
162
+ for i in range(len(self.layers)):
163
+ self.layers[i] = AssociativeLayerWrapper(
164
+ self.layers[i],
165
+ self.d_model,
166
+ self.num_mem_tokens,
167
+ self.d_mem,
168
+ correction,
169
+ info={'layer': i}
170
+ )
171
+ self.create_memory(num_mem_tokens)
172
+ self.wrap_pos = wrap_pos
173
+ if wrap_pos:
174
+ self.wrap_positional_embeddings(num_mem_tokens)
175
+
176
+ def generate_mode(self, is_on):
177
+ for layer in self.layers:
178
+ layer.generate_mode = is_on
179
+
180
+ def create_memory(self, num_mem_tokens):
181
+ self.num_mem_tokens = num_mem_tokens
182
+ embeddings = self.model.get_input_embeddings()
183
+ memory_dim = getattr(self.model.config, 'n_embd', self.model.config.hidden_size)
184
+ memory_weights = torch.randn((num_mem_tokens, memory_dim)) * embeddings.weight.data.std()
185
+ self.register_parameter('memory', torch.nn.Parameter(memory_weights, requires_grad=True))
186
+
187
+ def wrap_positional_embeddings(self, num_mem_tokens):
188
+ num_pos_embs, emb_dim = self.model.transformer.wpe.weight.shape
189
+ prev_embs = self.model.transformer.wpe.weight.detach()
190
+ self.model.transformer.wpe = torch.nn.Embedding(num_mem_tokens + num_pos_embs, emb_dim)
191
+
192
+ new_num_pos = num_pos_embs + num_mem_tokens
193
+ with torch.no_grad():
194
+ self.model.transformer.wpe.weight[:len(self.model.transformer.wpe.weight)-num_mem_tokens] = prev_embs
195
+ for layer in self.model.transformer.h:
196
+ layer.layer.attn.bias = torch.tril(torch.ones((new_num_pos, new_num_pos), dtype=torch.uint8)).view(
197
+ 1, 1, new_num_pos, new_num_pos
198
+ )
199
+
200
+ def set_memory(self, input_shape):
201
+ memory = self.memory.repeat(input_shape[0], 1, 1)
202
+ return memory
203
+
204
+ def zero_mem(self):
205
+ for layer in self.layers:
206
+ layer.zero_mem()
207
+ self.previous_input = None
208
+
209
+ def forward(self, input_ids, labels=None, labels_mask=None, zero_mem=False, **kwargs):
210
+ current_input_ids = input_ids.clone()
211
+ if self.attend_to_previous_input and self.previous_input is not None:
212
+ input_ids = torch.cat([self.previous_input, input_ids], dim=1)
213
+ if zero_mem:
214
+ self.zero_mem()
215
+
216
+
217
+ seg_kwargs = self.process_input(input_ids, **kwargs)
218
+
219
+ out = self.model(**seg_kwargs)
220
+
221
+ if self.attend_to_previous_input and self.previous_input is not None:
222
+ out['logits'] = out['logits'][:, self.previous_input.size(1):]
223
+ out = self.process_output(out, labels, labels_mask, **kwargs)
224
+
225
+ self.previous_input = current_input_ids
226
+ return out
227
+
228
+ def process_input(self, input_ids, **kwargs):
229
+ memory_state = self.set_memory(input_ids.shape)
230
+ seg_kwargs = dict(**kwargs)
231
+ inputs_embeds = kwargs.get('inputs_embeds')
232
+ if inputs_embeds is None:
233
+ inputs_embeds = self.model.get_input_embeddings()(input_ids)
234
+ inputs_embeds = torch.cat([inputs_embeds, memory_state], dim=1)
235
+
236
+ seg_kwargs['input_ids'] = None
237
+ seg_kwargs['inputs_embeds'] = inputs_embeds
238
+ if kwargs.get('attention_mask') is not None:
239
+ #seg_kwargs['attention_mask'] = self.pad_attention_mask(kwargs['attention_mask'], inputs_embeds.shape)
240
+ seg_kwargs['attention_mask'] = self.pad_attention_mask(kwargs['attention_mask'])
241
+ if kwargs.get('prev_attn_mask') is not None:
242
+ seg_kwargs['attention_mask'] = torch.cat([kwargs['prev_attn_mask'], seg_kwargs['attention_mask']], dim=-1)
243
+ if 'prev_attn_mask' in seg_kwargs.keys():
244
+ seg_kwargs.pop('prev_attn_mask')
245
+ seg_kwargs['output_hidden_states'] = True
246
+
247
+ if self.wrap_pos:
248
+ num_pos_embs = self.model.transformer.wpe.weight.shape[0]
249
+ ordinary_pos = torch.arange(0, input_ids.size(1), dtype=torch.long, device=input_ids.device)
250
+ write_pos = torch.arange(num_pos_embs - self.num_mem_tokens, num_pos_embs, dtype=torch.long, device=input_ids.device)
251
+ seg_kwargs['position_ids'] = torch.cat([
252
+ ordinary_pos,
253
+ write_pos
254
+ ]).long().unsqueeze(0)
255
+ return seg_kwargs
256
+
257
+ def pad_attention_mask(self, attention_mask):
258
+ if self.num_mem_tokens in {0, None}:
259
+ return attention_mask
260
+ else:
261
+ #mask = torch.ones(*shape[:2], dtype=torch.int64).to(attention_mask.device)
262
+ shape = list(attention_mask.shape)
263
+ shape[1] += self.num_mem_tokens
264
+ mask = torch.ones(*shape, dtype=torch.int64).to(attention_mask.device)
265
+ mask[:, :-self.num_mem_tokens] = attention_mask
266
+ return mask
267
+
268
+ def process_output(self, model_outputs, labels, labels_mask, **kwargs):
269
+ if self.num_mem_tokens not in {0, None}:
270
+ out = CausalLMOutputWithCrossAttentions()
271
+ out['logits'] = model_outputs.logits[:, :-self.num_mem_tokens]
272
+ if kwargs.get('output_hidden_states'):
273
+ out['hidden_states'] = [lh[:, :-self.num_mem_tokens] for lh in model_outputs.hidden_states]
274
+ if kwargs.get('output_attentions'):
275
+ out['attentions'] = model_outputs['attentions']
276
+ else:
277
+ out = model_outputs
278
+
279
+ if labels is not None:
280
+ ce_loss_fn = CrossEntropyLoss()
281
+ logits = out['logits'][..., :-1, :].contiguous()
282
+ flat_logits = logits.view(-1, logits.size(-1))
283
+ labels = labels[..., 1:].contiguous()
284
+ flat_labels = labels.view(-1)
285
+ if labels_mask is not None:
286
+ flat_mask = labels_mask[..., :-1].contiguous().view(-1)
287
+
288
+ flat_logits = flat_logits[flat_mask]
289
+ flat_labels = flat_labels[flat_mask]
290
+ ce_loss = ce_loss_fn(flat_logits, flat_labels)
291
+ out['ce_loss'] = ce_loss
292
+
293
+ if kwargs.get('use_cache') is not None:
294
+ out['past_key_values'] = model_outputs.past_key_values
295
+
296
+ return out
297
+
298
+ def generate(self, input_ids, attention_mask, zero_mem=False, **generate_kwargs):
299
+ if zero_mem:
300
+ self.zero_mem()
301
+
302
+
303
+ self.generate_mode(True)
304
+ seg_kwargs = self.process_input(input_ids, attention_mask=attention_mask)
305
+ out = self.model.generate(
306
+ inputs_embeds=seg_kwargs['inputs_embeds'][:, :-self.num_mem_tokens],
307
+ attention_mask=seg_kwargs['attention_mask'][:, :-self.num_mem_tokens],
308
+ **generate_kwargs
309
+ )
310
+ self.generate_mode(False)
311
+ return out
312
+
313
+
314
+ class AssociativeRecurrentWrapper(torch.nn.Module):
315
+ def __init__(self, memory_cell, **rmt_kwargs):
316
+ super().__init__()
317
+
318
+ self.memory_cell = memory_cell
319
+ self.rmt_config = rmt_kwargs
320
+
321
+ def forward(self,
322
+ input_ids,
323
+ labels=None,
324
+ labels_mask=None,
325
+ inputs_embeds=None,
326
+ attention_mask=None,
327
+ output_attentions=None,
328
+ output_hidden_states=None,
329
+ input_segmented=False,
330
+ sliding_window=False,
331
+ ):
332
+ attend_to_previous_input = self.rmt_config['attend_to_previous_input'] if 'attend_to_previous_input' in self.rmt_config else False
333
+ if input_segmented:
334
+ n_segs = input_ids.shape[1] if not (input_ids is None) else inputs_embeds.shape[1]
335
+ segmented = [dict(
336
+ input_ids=input_ids[:, i] if not (input_ids is None) else None,
337
+ inputs_embeds=inputs_embeds[:, i] if not (inputs_embeds is None) else None,
338
+ attention_mask=attention_mask[:, i],
339
+ labels=labels[:, i] if not (labels is None) else None,
340
+ labels_mask=labels_mask[:, i] if not (labels_mask is None) else None,
341
+ ) for i in range(n_segs)]
342
+ labels = torch.cat([labels[:, i] for i in range(n_segs)], dim=1)
343
+ if labels_mask is not None:
344
+ labels_mask = torch.cat([labels_mask[:, i] for i in range(n_segs)], dim=1)
345
+ else:
346
+ segmented = self.segment(input_ids=input_ids, inputs_embeds=inputs_embeds, attention_mask=attention_mask, labels=labels, labels_mask=labels_mask)
347
+ cell_outputs = []
348
+ past_key_values = None
349
+ num_mem_tokens = self.memory_cell.num_mem_tokens
350
+ prev_attn_mask = None
351
+ self.memory_cell.zero_mem()
352
+ for seg_num, segment in enumerate(segmented):
353
+ seg_len = segment['input_ids'].size(-1)
354
+ cell_out = self.memory_cell(**segment,
355
+ output_hidden_states=True,
356
+ use_cache=sliding_window,
357
+ past_key_values=past_key_values,
358
+ prev_attn_mask=prev_attn_mask,
359
+ zero_mem=False
360
+ )
361
+ if sliding_window or attend_to_previous_input:
362
+ prev_attn_mask = segment['attention_mask'] * torch.triu(torch.ones_like(segment['attention_mask']))
363
+ if sliding_window:
364
+ past_key_values = [
365
+ [
366
+ k_or_v[..., -(num_mem_tokens+seg_len):k_or_v.size(-2)-num_mem_tokens, :].detach()
367
+ for k_or_v in seg_kv
368
+ ]
369
+ for seg_kv in cell_out['past_key_values']
370
+ ]
371
+ cell_outputs.append(cell_out)
372
+ self.memory_cell.zero_mem()
373
+
374
+
375
+ out = self.process_outputs(cell_outputs, labels=labels,
376
+ labels_mask=labels_mask,
377
+ output_attentions=output_attentions,
378
+ output_hidden_states=output_hidden_states)
379
+ return out
380
+
381
+ def segment(self, **kwargs):
382
+ segments = []
383
+ for k, tensor in kwargs.items():
384
+ if tensor is not None:
385
+ k_segments = self.split_tensor(tensor)
386
+ for s, k_seg in enumerate(k_segments):
387
+ if s < len(segments):
388
+ segments[s][k] = k_seg
389
+ else:
390
+ segments.append({k: k_seg})
391
+
392
+ return segments
393
+
394
+ def split_tensor(self, tensor):
395
+ align = self.rmt_config.get('segment_alignment')
396
+ segment_size = self.rmt_config.get('segment_size')
397
+ if align in {'left', None}:
398
+ split_inds = list(range(0, tensor.shape[1], segment_size)) + [tensor.shape[1]]
399
+ segments = [tensor[:, start:end] for (start, end) in zip(split_inds, split_inds[1:])]
400
+ elif align in {'right', None}:
401
+ split_inds = (list(range(tensor.shape[1], 0, -segment_size)) + [0])[::-1]
402
+ segments = [tensor[:, start:end] for (start, end) in zip(split_inds, split_inds[1:])]
403
+ elif align == 'center':
404
+ n_seg = math.ceil(tensor.shape[1] / segment_size)
405
+ segments = torch.chunk(tensor, n_seg, dim=1)
406
+ else:
407
+ raise NotImplementedError
408
+ return segments
409
+
410
+ def process_outputs(self, cell_outputs, **kwargs):
411
+ out = CausalLMOutputWithCrossAttentions()
412
+ full_logits = torch.cat([o.logits for o in cell_outputs], dim=1)
413
+ full_hidden_states = tuple([torch.cat(layer_hs, dim=1) for layer_hs in zip(*[o.hidden_states for o in cell_outputs])])
414
+
415
+ labels = kwargs.get('labels')
416
+ if labels is not None:
417
+ shift_labels = labels[..., 1:].contiguous()
418
+ shift_logits = full_logits[..., :-1, :].contiguous()
419
+ flat_labels = shift_labels.view(-1)
420
+ flat_logits = shift_logits.view(-1, shift_logits.size(-1))
421
+
422
+ loss_fct = CrossEntropyLoss()
423
+ labels_mask = kwargs.get('labels_mask')
424
+ if labels_mask is not None:
425
+ shift_mask = labels_mask[..., :-1].contiguous()
426
+
427
+ flat_labels = flat_labels[shift_mask.view(-1)]
428
+ flat_logits = flat_logits[shift_mask.view(-1)]
429
+
430
+ out['loss'] = loss_fct(flat_logits, flat_labels)
431
+ else:
432
+ out['loss'] = 0
433
+
434
+ if self.rmt_config.get("return_all_logits", False):
435
+ out['ce_loss'] = out['loss']
436
+
437
+ out['logits'] = full_logits
438
+ segment_keys = ['loss', 'logits']
439
+ if kwargs.get('output_attentions'):
440
+ segment_keys.append('attentions')
441
+ if kwargs.get('output_hidden_states'):
442
+ segment_keys.append('hidden_states')
443
+ out['hidden_states'] = full_hidden_states
444
+
445
+ if self.rmt_config.get("return_all_logits", False):
446
+ for seg_num, o in enumerate(cell_outputs):
447
+ for key, value in o.items():
448
+ if any([sk in key for sk in segment_keys]):
449
+ out[f'{key}_{seg_num}'] = value
450
+ return out
451
+
452
+ def manage_gradients(self, memory_state, seg_num):
453
+ k2, max_n_segments = self.rmt_config.get('k2'), self.rmt_config.get('max_n_segments')
454
+ if seg_num == 0 \
455
+ or k2 in {-1, None} \
456
+ or seg_num + k2 > max_n_segments:
457
+ return True
458
+
459
+ memory_state = memory_state.detach()
460
+ return False
461
+
462
+ def generate(self, input_ids, attention_mask, **generate_kwargs):
463
+ self.memory_cell.zero_mem()
464
+ segmented = self.segment(input_ids=input_ids, attention_mask=attention_mask)
465
+
466
+ for seg_num, segment in enumerate(segmented[:-1]):
467
+ cell_out = self.memory_cell(**segment, output_hidden_states=True, zero_mem=False)
468
+
469
+ final_segment = segmented[-1]
470
+ out = self.memory_cell.generate(**final_segment, zero_mem=False, **generate_kwargs)
471
+ self.memory_cell.zero_mem()
472
+ return out
473
+
474
+ def gradient_checkpointing_enable(self, *args, **kwargs):
475
+ # doesn't supported for ARMT
476
+ self.memory_cell.model.gradient_checkpointing_enable(*args, **kwargs)
language_modeling_amt.py ADDED
@@ -0,0 +1,1157 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # copy of https://github.com/RodkinIvan/associative-recurrent-memory-transformer/blob/llama_armt/modeling_amt/language_modeling.py
2
+ # with small changes for compatibility
3
+ import math
4
+ import torch
5
+ from torch.nn import CrossEntropyLoss
6
+ from transformers.modeling_outputs import CausalLMOutputWithCrossAttentions
7
+ from transformers.cache_utils import Cache, DynamicCache
8
+ from torch.nn.functional import relu as r
9
+ import torch.nn.functional as F
10
+ from munch import Munch
11
+ import os
12
+
13
+
14
+ from modeling_amt.act_utils import ACT_basic, gen_timing_signal
15
+ # from baselines.rwkv.language_modeling import RWKVModel
16
+
17
+ def dpfp(x, nu=1):
18
+ x = torch.cat([r(x), r(-x)], dim=-1)
19
+ x_rolled = torch.cat([x.roll(shifts=j, dims=-1)
20
+ for j in range(1,nu+1)], dim=-1)
21
+ x_repeat = torch.cat([x] * nu, dim=-1)
22
+ return x_repeat * x_rolled
23
+
24
+ class DPFP:
25
+ def __init__(self, nu):
26
+ self.nu = nu
27
+
28
+ def __call__(self, x):
29
+ nu = self.nu
30
+ x = torch.cat([r(x), r(-x)], dim=-1)
31
+ x_rolled = torch.cat([x.roll(shifts=j, dims=-1) for j in range(1,nu+1)], dim=-1)
32
+ x_repeat = torch.cat([x] * nu, dim=-1)
33
+ return x_repeat * x_rolled
34
+
35
+
36
+ def attn_mask_to_4d(attn_mask, upper, query_len):
37
+ if attn_mask is None:
38
+ return None
39
+ seg_len = attn_mask.size(-1)
40
+ if upper:
41
+ tri = torch.triu(torch.ones(query_len, seg_len))
42
+ else:
43
+ tri = torch.tril(torch.ones(query_len, seg_len))
44
+
45
+ mask = torch.einsum('bj,ij->bij', attn_mask, tri.to(attn_mask.device))
46
+ mask = mask.unsqueeze(1)
47
+ return mask
48
+
49
+ def invert_attn_mask(attn_mask, dtype):
50
+ min_dtype = torch.finfo(dtype).min
51
+ new_mask = (1.0 - attn_mask) * min_dtype
52
+ return new_mask
53
+
54
+
55
+ class AssociativeLayerWrapper(torch.nn.Module):
56
+
57
+ def __init__(self, layer, d_model, num_mem_tokens, d_mem, n_heads=1, correction=True, info=None, use_denom=True, gating=False, compress_mem=0) -> None:
58
+ super().__init__()
59
+ self.info = info
60
+ self.seg_num = 0
61
+ self.d_model = d_model
62
+ self.num_mem_tokens = num_mem_tokens
63
+ self.d_mem = d_mem
64
+ self.n_heads = n_heads
65
+ self.gating = gating
66
+ self.compress_mem = compress_mem
67
+ nu = 3
68
+ self.d_key = 2 * nu * d_mem
69
+
70
+ assert self.d_mem % n_heads == 0 and self.d_model % n_heads == 0
71
+
72
+ self.phi = DPFP(nu)
73
+ # self.d_key = d_mem
74
+ # self.phi = torch.nn.Identity()
75
+
76
+ self.use_denom = use_denom
77
+
78
+ self.W_mq = torch.nn.Linear(d_model, d_mem, bias=False)
79
+ # torch.nn.init.zeros_(self.W_mq.weight)
80
+ self.W_mk = torch.nn.Linear(d_model, d_mem, bias=False)
81
+ if self.compress_mem != 0:
82
+ self.W_mv_in = torch.nn.Linear(d_model, self.compress_mem, bias=False)
83
+ self.W_mv_out = torch.nn.Linear(self.compress_mem, d_model, bias=False)
84
+ torch.nn.init.zeros_(self.W_mv_in.weight)
85
+ torch.nn.init.zeros_(self.W_mv_out.weight)
86
+ else:
87
+ self.W_mv = torch.nn.Linear(d_model, d_model, bias=False)
88
+ torch.nn.init.zeros_(self.W_mv.weight)
89
+ if gating:
90
+ self.W_mb = torch.nn.Linear(d_model, d_model)
91
+ else:
92
+ self.W_mb = torch.nn.Linear(d_model, n_heads)
93
+
94
+ self.W_mem = torch.zeros(1, n_heads ,self.d_key // n_heads, d_model // n_heads)
95
+ self.W_mem.requires_grad_(False)
96
+ if self.use_denom:
97
+ self.z = torch.zeros(1, n_heads, self.d_key // n_heads)
98
+ self.z.requires_grad_(False)
99
+
100
+ # self.ln = torch.nn.LayerNorm(d_model)
101
+
102
+ self.zero_mem()
103
+
104
+ self.layer = layer
105
+
106
+ self.generate_mode = False
107
+ self.first_seg = True
108
+ self.correction = correction
109
+
110
+
111
+ def _to_heads(self, x):
112
+ bsz, seq_len, d_model = x.shape
113
+ x = x.reshape(bsz, seq_len, self.n_heads, d_model // self.n_heads)
114
+ x = x.permute(0, 2, 1, 3)
115
+ return x
116
+
117
+ def _from_heads(self, x):
118
+ bsz, n_heads, seq_len, d_head = x.shape
119
+ x = x.permute(0, 2, 1, 3).reshape(bsz, seq_len, n_heads * d_head)
120
+ return x
121
+ def associate(self, hidden_states):
122
+ bsz, seq_len, d_model = hidden_states.shape
123
+
124
+ self.W_mem = self.W_mem.to(hidden_states.device)
125
+ if self.use_denom:
126
+ self.z = self.z.to(hidden_states.device)
127
+
128
+ q = self._to_heads(self.W_mq(hidden_states))
129
+ mq = self.phi(q) # (bsz, n_heads, seq_len, 2 * d_head * nu)
130
+ mq = F.normalize(mq, dim=-1, p=2.0)
131
+ # crutch for dataparallel
132
+ # mq += 0 * self.W_mb(hidden_states).sum() * self.W_mk(hidden_states).sum() * self.W_mv(hidden_states).sum()
133
+
134
+ num = torch.einsum('ihjk,ihkt->ihjt', mq, self.W_mem)
135
+ if self.use_denom:
136
+ denom = torch.einsum("ihk,ihjk->ihj", self.z, mq)[..., None] + 1e-5
137
+ hidden_states = num / denom # (bsz, n_heads, seq_len, d_model // n_heads)
138
+ else:
139
+ hidden_states = num
140
+ hidden_states = self._from_heads(hidden_states)
141
+ return hidden_states
142
+
143
+ def forward(self, hidden_states, *args, **kwargs):
144
+ if not self.first_seg:
145
+ hidden_states = self.associate(
146
+ # self.ln(
147
+ hidden_states
148
+ # )
149
+ ) + hidden_states
150
+ out = self.layer(hidden_states, *args, **kwargs)
151
+ if not self.generate_mode:
152
+ mem_tokens = out[0][:, -self.num_mem_tokens:]
153
+ # mem_tokens = out[0]
154
+ self.update_mem(mem_tokens)
155
+ self.first_seg = False
156
+ return out
157
+
158
+ def forward_no_update(self, hidden_states, *args, **kwargs):
159
+ if not self.first_seg:
160
+ hidden_states = self.associate(
161
+ # self.ln(
162
+ hidden_states
163
+ # )
164
+ ) + hidden_states
165
+ out = self.layer(hidden_states, *args, **kwargs)
166
+ return out
167
+
168
+ def update_mem(self, mem_tokens):
169
+
170
+ self.W_mem = self.W_mem.to(mem_tokens.device)
171
+ if self.use_denom:
172
+ self.z = self.z.to(mem_tokens.device)
173
+ k = self._to_heads(self.W_mk(mem_tokens))
174
+ mk = self.phi(k)
175
+ mk = F.normalize(mk, dim=-1, p=2.0)
176
+
177
+ if self.compress_mem != 0:
178
+ new_mv = self.W_mv_in(mem_tokens)
179
+ new_mv = self._to_heads(self.W_mv_out(new_mv))
180
+ else:
181
+ new_mv = self._to_heads(self.W_mv(mem_tokens)) # (bsz, n_heads, num_mem_tokens, d_model)
182
+ if not self.first_seg:
183
+ num = torch.einsum('ihjk,ihkt->ihjt', mk, self.W_mem)
184
+ if self.use_denom:
185
+ denom = torch.einsum("ihj,ihkj->ihk", self.z, mk)[..., None] + 1e-5
186
+ prev_mv = num / denom
187
+ if self.correction:
188
+ new_info_coef = (1 - denom / (torch.linalg.norm(mk, dim=-1) ** 2)[..., None])
189
+ new_info_coef = torch.clip(new_info_coef, 0, 1).detach()
190
+ else:
191
+ new_info_coef = 1
192
+ else:
193
+ prev_mv = num
194
+ else:
195
+ prev_mv = torch.zeros_like(new_mv, device=new_mv.device)
196
+ new_info_coef = 1
197
+
198
+ # wandb.log({f"gamma_{self.info['layer']}": new_info_coef.mean(dim=1).item() if isinstance(new_info_coef, torch.Tensor) else 1}, step=self.seg_num)
199
+ mv = new_mv - prev_mv
200
+
201
+ # new_norm = torch.linalg.norm(new_mv, dim=-1)
202
+ # old_norm = torch.linalg.norm(prev_mv, dim=-1)
203
+ # new_info_coef = torch.clip(1 - old_norm / (new_norm + 1e-5), -10, 10)[..., None].detach()
204
+ # new_info_coef = 1 - denom
205
+
206
+ mb = self._to_heads(torch.sigmoid(self.W_mb(mem_tokens)))
207
+
208
+ einop = f"ihjk,ihjt,ihj{'t' if self.gating else 'x'}->ihkt"
209
+ associations = torch.einsum(einop, mk, mv, mb) # (bsz, n_heads, d_mem, d_model)
210
+
211
+ self.W_mem = self.W_mem + associations
212
+
213
+ if self.use_denom:
214
+ self.z = self.z + (new_info_coef*mk).sum(dim=-2)
215
+ # self.z = self.z + (new_info_coef*mb[..., None]*mk).sum(dim=1)
216
+ self.seg_num += 1
217
+
218
+ def freeze_mem(self):
219
+ self.W_mb.weight.requires_grad = False
220
+ self.W_mb.bias.requires_grad = False
221
+
222
+ self.W_mq.weight.requires_grad = False
223
+ self.W_mk.weight.requires_grad = False
224
+ if self.compress_mem != 0:
225
+ self.W_mv_in.weight.requires_grad = False
226
+ self.W_mv_out.weight.requires_grad = False
227
+ else:
228
+ self.W_mv.weight.requires_grad = False
229
+
230
+ def zero_mem(self):
231
+ self.first_seg = True
232
+ self.W_mem = torch.zeros(1, self.n_heads, self.d_key // self.n_heads, self.d_model // self.n_heads).to(next(self.parameters()).dtype)
233
+ if self.use_denom:
234
+ self.z = torch.zeros(1, self.n_heads, self.d_key // self.n_heads).to(next(self.parameters()).dtype)
235
+ self.seg_num = 0
236
+
237
+
238
+
239
+ class AdaptiveAssociativeLayerWrapper(AssociativeLayerWrapper):
240
+ def __init__(self,
241
+ layer,
242
+ d_model,
243
+ num_mem_tokens,
244
+ d_mem,
245
+ max_hop,
246
+ n_heads=1,
247
+ correction=True,
248
+ info=None,
249
+ use_denom=True,
250
+ gating=False,
251
+
252
+ ) -> None:
253
+ super().__init__(layer, d_model, num_mem_tokens, d_mem, n_heads, correction, info, use_denom, gating)
254
+ self.act = ACT_basic(d_model)
255
+ self.depth = max_hop
256
+ self.max_length = 1024
257
+
258
+ self.timing_signal = gen_timing_signal(self.max_length, d_model)
259
+ ## for t
260
+ self.position_signal = gen_timing_signal(self.depth, d_model)
261
+
262
+ self.remainders = torch.zeros(1,)
263
+ self.n_updates = torch.zeros(1,)
264
+ self.segments_passed = torch.zeros(1,)
265
+
266
+ def associate(self, hidden_states):
267
+ self.remainders = self.remainders.to(hidden_states.device)
268
+ self.n_updates = self.n_updates.to(hidden_states.device)
269
+ self.segments_passed = self.segments_passed.to(hidden_states.device)
270
+ out, (remainders, n_updates) = self.act(
271
+ state=hidden_states,
272
+ inputs=hidden_states,
273
+ fn=super().associate,
274
+ time_enc=self.timing_signal,
275
+ pos_enc=self.position_signal,
276
+ max_hop=self.depth
277
+ )
278
+
279
+ self.remainders = self.remainders + remainders # 1 - \sum(h_i); L' = L + tau * mean(remainders)
280
+ self.n_updates = self.n_updates + n_updates
281
+ self.segments_passed = self.segments_passed + 1
282
+ return out
283
+
284
+ def zero_mem(self):
285
+ self.remainders = torch.zeros(1,)
286
+ self.n_updates = torch.zeros(1,)
287
+ self.segments_passed = torch.zeros(1,)
288
+ return super().zero_mem()
289
+
290
+
291
+
292
+
293
+ class AdaptiveAssociativeLayerWrapper2(AssociativeLayerWrapper):
294
+ def __init__(self,
295
+ layer,
296
+ d_model,
297
+ num_mem_tokens,
298
+ d_mem,
299
+ max_hop,
300
+ n_heads=1,
301
+ correction=True,
302
+ info=None,
303
+ use_denom=True,
304
+ gating=False,
305
+
306
+ ) -> None:
307
+ super().__init__(layer, d_model, num_mem_tokens, d_mem, n_heads, correction, info, use_denom, gating)
308
+ self.act = ACT_basic(d_model)
309
+ self.depth = max_hop
310
+ self.max_length = 1024
311
+
312
+ self.timing_signal = gen_timing_signal(self.max_length, d_model)
313
+ ## for t
314
+ self.position_signal = gen_timing_signal(self.depth, d_model)
315
+
316
+ self.remainders = torch.zeros(1,)
317
+ self.n_updates = torch.zeros(1,)
318
+ self.segments_passed = torch.zeros(1,)
319
+
320
+ def forward(self, hidden_states, *args, **kwargs):
321
+ self.remainders = self.remainders.to(hidden_states.device)
322
+ self.n_updates = self.n_updates.to(hidden_states.device)
323
+ self.segments_passed = self.segments_passed.to(hidden_states.device)
324
+
325
+ fwd = super().forward_no_update
326
+ out, (remainders, n_updates) = self.act(
327
+ *args,
328
+ state=hidden_states,
329
+ inputs=hidden_states,
330
+ fn=fwd,
331
+ time_enc=self.timing_signal,
332
+ pos_enc=self.position_signal,
333
+ max_hop=self.depth,
334
+ **kwargs
335
+ )
336
+ if not self.generate_mode:
337
+ mem_tokens = out[0][:, -self.num_mem_tokens:]
338
+ # mem_tokens = out[0]
339
+ self.update_mem(mem_tokens)
340
+ self.first_seg = False
341
+ self.remainders = self.remainders + remainders # 1 - \sum(h_i); L' = L + tau * mean(reminders)
342
+ self.n_updates = self.n_updates + n_updates
343
+ self.segments_passed = self.segments_passed + 1
344
+ return out
345
+
346
+
347
+ def zero_mem(self):
348
+ self.remainders = torch.zeros(1,)
349
+ self.n_updates = torch.zeros(1,)
350
+ self.segments_passed = torch.zeros(1,)
351
+ return super().zero_mem()
352
+
353
+
354
+ class AssociativeMemoryCell(torch.nn.Module):
355
+ def __init__(self,
356
+ base_model,
357
+ num_mem_tokens,
358
+ d_mem,
359
+ layers_attr: str = 'model.layers',
360
+ wrap_pos=False,
361
+ correction=True,
362
+ n_heads=1,
363
+ use_denom=True,
364
+ gating=False,
365
+ freeze_mem=False,
366
+ act_on=False,
367
+ max_hop=4,
368
+ act_type='associative',
369
+ attend_to_previous_input=False,
370
+ use_sink=False,
371
+ use_lora=False,
372
+ compress_mem=0,
373
+ ):
374
+ super().__init__()
375
+ self.model = base_model
376
+ self.attend_to_previous_input = attend_to_previous_input
377
+ self.previous_input = None
378
+ self.use_sink = use_sink
379
+
380
+ self.RWKV_ARMT = False #isinstance(self.model, RWKVModel)
381
+
382
+ self.num_mem_tokens = num_mem_tokens
383
+ self.d_mem = d_mem
384
+ self.d_model = base_model.get_input_embeddings().embedding_dim
385
+ self.W_mem = []
386
+ if use_lora:
387
+ # LoRA case
388
+ self.layers = self.model.model
389
+ else:
390
+ self.layers = self.model
391
+
392
+ self.layers_attrs = layers_attr.split('.')
393
+ for i, attr in enumerate(self.layers_attrs):
394
+ self.layers = getattr(self.layers, attr)
395
+
396
+ for i in range(len(self.layers)):
397
+ kw = dict(
398
+ layer=self.layers[i],
399
+ d_model=self.d_model,
400
+ num_mem_tokens=self.num_mem_tokens,
401
+ d_mem=self.d_mem,
402
+ correction=correction,
403
+ info={'layer': i},
404
+ n_heads=n_heads,
405
+ use_denom=use_denom,
406
+ gating=gating,
407
+ compress_mem=compress_mem
408
+ )
409
+ if act_on:
410
+ kw['max_hop'] = max_hop
411
+ if not act_on:
412
+ self.layers[i] = AssociativeLayerWrapper(**kw)
413
+ elif act_type == 'associative':
414
+ self.layers[i] = AdaptiveAssociativeLayerWrapper(**kw)
415
+ elif act_type == 'layer':
416
+ self.layers[i] = AdaptiveAssociativeLayerWrapper2(**kw)
417
+ else:
418
+ raise f'Unknown ACT type: {act_type}'
419
+ self.create_memory(num_mem_tokens)
420
+ self.wrap_pos = wrap_pos
421
+ self.act_on = act_on
422
+ if wrap_pos:
423
+ self.wrap_positional_embeddings(num_mem_tokens)
424
+
425
+ if freeze_mem:
426
+ for layer in self.layers:
427
+ layer.freeze_mem()
428
+
429
+ def generate_mode(self, is_on):
430
+ for layer in self.layers:
431
+ layer.generate_mode = is_on
432
+
433
+ def create_memory(self, num_mem_tokens):
434
+ self.num_mem_tokens = num_mem_tokens
435
+ embeddings = self.model.get_input_embeddings()
436
+ memory_dim = getattr(self.model.config, 'n_embd', self.model.config.hidden_size)
437
+ memory_weights = torch.randn((num_mem_tokens, memory_dim), device=embeddings.weight.data.device) * embeddings.weight.data.std()
438
+ self.register_parameter('memory', torch.nn.Parameter(memory_weights, requires_grad=True))
439
+ if self.use_sink:
440
+ self.sink = torch.nn.Parameter(torch.randn((1, memory_dim), device=embeddings.weight.data.device), requires_grad=True)
441
+
442
+
443
+ def wrap_positional_embeddings(self, num_mem_tokens):
444
+ num_pos_embs, emb_dim = self.model.transformer.wpe.weight.shape
445
+ prev_embs = self.model.transformer.wpe.weight.detach()
446
+ self.model.transformer.wpe = torch.nn.Embedding(num_mem_tokens + num_pos_embs, emb_dim)
447
+
448
+ new_num_pos = num_pos_embs + num_mem_tokens
449
+ with torch.no_grad():
450
+ self.model.transformer.wpe.weight[:len(self.model.transformer.wpe.weight)-num_mem_tokens] = prev_embs
451
+ for layer in self.model.transformer.h:
452
+ layer.layer.attn.bias = torch.tril(torch.ones((new_num_pos, new_num_pos), dtype=torch.uint8)).view(
453
+ 1, 1, new_num_pos, new_num_pos
454
+ )
455
+
456
+ def set_memory(self, input_shape):
457
+ memory = self.memory.repeat(input_shape[0], 1, 1)
458
+ if self.use_sink:
459
+ sink = self.sink.repeat(input_shape[0], 1, 1)
460
+ else:
461
+ sink = None
462
+ return memory, sink
463
+
464
+ def zero_mem(self):
465
+ for layer in self.layers:
466
+ layer.zero_mem()
467
+ pass
468
+ self.previous_input = None
469
+
470
+ def forward(self, input_ids, labels=None, labels_mask=None, zero_mem=False, **kwargs):
471
+ current_input_ids = input_ids.clone()
472
+ if self.attend_to_previous_input and self.previous_input is not None:
473
+ input_ids = torch.cat([self.previous_input, input_ids], dim=1)
474
+
475
+ if zero_mem:
476
+ self.zero_mem()
477
+ seg_kwargs = self.process_input(input_ids, **kwargs)
478
+
479
+ if self.RWKV_ARMT and not self.layers[0].generate_mode:
480
+ input1 = dict()
481
+ input2 = dict()
482
+ for item in seg_kwargs:
483
+ if isinstance(seg_kwargs[item], torch.Tensor):
484
+ # if False:
485
+ input1[item] = seg_kwargs[item][:, :-self.num_mem_tokens]
486
+ input2[item] = seg_kwargs[item][:, -self.num_mem_tokens:]
487
+ else:
488
+ input1[item] = seg_kwargs[item]
489
+ input2[item] = seg_kwargs[item]
490
+
491
+ self.generate_mode(True)
492
+ out = self.model(**input1)
493
+ self.generate_mode(False)
494
+ state_tmp = tuple([torch.clone(state) for state in out['state']])
495
+ out = Munch({k: torch.clone(t) if isinstance(t, torch.Tensor) else t for k, t in out.items()})
496
+ input2['state'] = out['state']
497
+ _ = self.model(**input2)
498
+ out['state'] = state_tmp
499
+ # out['state'] = out2['state']
500
+ # out = self.model(**seg_kwargs)
501
+ # out['logits'] = out['logits'][:, :-self.num_mem_tokens]
502
+ else:
503
+ out = self.model(**seg_kwargs)
504
+
505
+ if self.attend_to_previous_input and self.previous_input is not None:
506
+ out['logits'] = out['logits'][:, self.previous_input.size(1):]
507
+ out = self.process_output(out, labels, labels_mask, **kwargs)
508
+ self.previous_input = current_input_ids
509
+ return out
510
+
511
+ def process_input(self, input_ids, **kwargs):
512
+ memory_state, sink = self.set_memory(input_ids.shape)
513
+ seg_kwargs = dict(**kwargs)
514
+ inputs_embeds = kwargs.get('inputs_embeds')
515
+ if inputs_embeds is None:
516
+ inputs_embeds = self.model.get_input_embeddings()(input_ids)
517
+ if self.use_sink:
518
+ inputs_embeds = torch.cat([sink, inputs_embeds, memory_state], dim=1)
519
+ else:
520
+ inputs_embeds = torch.cat([inputs_embeds, memory_state], dim=1)
521
+
522
+ seg_kwargs['input_ids'] = None
523
+ seg_kwargs['inputs_embeds'] = inputs_embeds
524
+ if kwargs.get('attention_mask') is not None:
525
+ #print(kwargs['attention_mask'].shape)
526
+ seg_kwargs['attention_mask'] = self.pad_attention_mask(kwargs['attention_mask'], dtype=inputs_embeds.dtype)
527
+ if kwargs.get('prev_attn_mask') is not None:
528
+ #print(kwargs['prev_attn_mask'].shape)
529
+ prev_seg_attn_mask = self.pad_prev_seg_attn_mask(kwargs['prev_attn_mask'], dtype=inputs_embeds.dtype)
530
+ #print(prev_seg_attn_mask.shape, seg_kwargs['attention_mask'].shape, seg_kwargs['inputs_embeds'].shape)
531
+ seg_kwargs['attention_mask'] = torch.cat([prev_seg_attn_mask, seg_kwargs['attention_mask']], dim=-1)
532
+ if 'prev_attn_mask' in seg_kwargs:
533
+ seg_kwargs.pop('prev_attn_mask')
534
+ seg_kwargs['output_hidden_states'] = True
535
+
536
+ if self.wrap_pos:
537
+ num_pos_embs = self.model.transformer.wpe.weight.shape[0]
538
+ ordinary_pos = torch.arange(0, input_ids.size(1), dtype=torch.long, device=input_ids.device)
539
+ write_pos = torch.arange(num_pos_embs - self.num_mem_tokens, num_pos_embs, dtype=torch.long, device=input_ids.device)
540
+ seg_kwargs['position_ids'] = torch.cat([
541
+ ordinary_pos,
542
+ write_pos
543
+ ]).long().unsqueeze(0)
544
+ return seg_kwargs
545
+
546
+ def convert_to_infinity_attn_mask(self, attn_mask, dtype):
547
+ min_dtype = torch.finfo(dtype).min
548
+ new_mask = (1.0 - attn_mask) * min_dtype
549
+ return new_mask
550
+
551
+ def pad_attention_mask(self, attention_mask, dtype=float):
552
+ if self.num_mem_tokens in {0, None}:
553
+ return attention_mask
554
+ else:
555
+ shape = list(attention_mask.shape)
556
+ if len(shape) == 4:
557
+
558
+ shape[-1] += self.num_mem_tokens + self.use_sink
559
+ shape[-2] += self.num_mem_tokens + self.use_sink
560
+ mask = torch.ones(*shape, dtype=dtype).to(attention_mask.device)
561
+ mask[..., int(self.use_sink):-self.num_mem_tokens, int(self.use_sink):-self.num_mem_tokens] = attention_mask
562
+ if self.use_sink:
563
+ mask[..., 0, 1:] = 0
564
+ mask[..., :-self.num_mem_tokens, -self.num_mem_tokens:] = 0
565
+ # mask = torch.tril(mask)
566
+ if not os.environ.get("NOT_INVERT_ATTN_MASK"):
567
+ mask = invert_attn_mask(mask, dtype)
568
+ else:
569
+ shape[-1] += self.num_mem_tokens + self.use_sink
570
+ mask = torch.ones(*shape, dtype=dtype).to(attention_mask.device)
571
+ mask[..., int(self.use_sink):-self.num_mem_tokens] = attention_mask
572
+ return mask.to(dtype)
573
+
574
+ def pad_prev_seg_attn_mask(self, prev_seg_attn_mask, dtype=float):
575
+ if self.num_mem_tokens in {0, None}:
576
+ return prev_seg_attn_mask
577
+ else:
578
+ shape = list(prev_seg_attn_mask.shape)
579
+ if len(shape) == 4:
580
+ shape[-2] += self.num_mem_tokens + self.use_sink
581
+ mask = torch.ones(*shape, dtype=dtype).to(prev_seg_attn_mask.device)
582
+ mask[..., int(self.use_sink):-self.num_mem_tokens, :] = prev_seg_attn_mask
583
+ if self.use_sink:
584
+ mask[..., 0, :] = 0
585
+ if not os.environ.get("NOT_INVERT_ATTN_MASK"):
586
+ mask = invert_attn_mask(mask, dtype)
587
+ else:
588
+ mask = prev_seg_attn_mask
589
+ return mask.to(dtype)
590
+
591
+ def process_output(self, model_outputs, labels, labels_mask, **kwargs):
592
+
593
+
594
+ if (self.num_mem_tokens not in {0, None}) and not self.RWKV_ARMT:
595
+ out = CausalLMOutputWithCrossAttentions()
596
+ out['logits'] = model_outputs.logits[:, int(self.use_sink):-self.num_mem_tokens]
597
+ if kwargs.get('output_hidden_states'):
598
+ out['hidden_states'] = [lh[:, int(self.use_sink):-self.num_mem_tokens] for lh in model_outputs.hidden_states]
599
+ if kwargs.get('output_attentions'):
600
+ out['attentions'] = model_outputs['attentions']
601
+ else:
602
+ out = model_outputs
603
+
604
+ if labels is not None:
605
+ ce_loss_fn = CrossEntropyLoss()
606
+ logits = out['logits'][..., :-1, :].contiguous()
607
+ flat_logits = logits.view(-1, logits.size(-1))
608
+ labels = labels[..., 1:].contiguous()
609
+ flat_labels = labels.view(-1)
610
+ if labels_mask is not None:
611
+ flat_mask = labels_mask[..., :-1].contiguous().view(-1)
612
+ flat_logits = flat_logits[flat_mask]
613
+ flat_labels = flat_labels[flat_mask]
614
+ ce_loss = ce_loss_fn(flat_logits, flat_labels)
615
+ out['ce_loss'] = ce_loss
616
+
617
+ if kwargs.get('use_cache', False):
618
+ out['past_key_values'] = model_outputs.past_key_values
619
+
620
+ return out
621
+
622
+ def generate(self, input_ids, attention_mask, prev_attn_mask=None, use_cache=False, past_key_values=None, zero_mem=False, **generate_kwargs):
623
+ if zero_mem:
624
+ self.zero_mem()
625
+
626
+
627
+ self.generate_mode(True)
628
+ inp_kwargs = {
629
+ "attention_mask": attention_mask,
630
+ #"prev_attn_mask": prev_attn_mask,
631
+ #"use_cache": use_cache,
632
+ #"past_key_values": past_key_values,
633
+ }
634
+ seg_kwargs = self.process_input(input_ids, **inp_kwargs)
635
+ #print(seg_kwargs)
636
+ #print(seg_kwargs["inputs_embeds"].shape)
637
+ #print(seg_kwargs["attention_mask"].shape)
638
+ #print(smth)
639
+ out = self.model.generate(
640
+ inputs_embeds=seg_kwargs['inputs_embeds'][:, :-self.num_mem_tokens],
641
+ attention_mask=seg_kwargs['attention_mask'][:, :-self.num_mem_tokens],
642
+ **generate_kwargs
643
+ )
644
+ #print(smth)
645
+ self.generate_mode(False)
646
+ return out
647
+
648
+ def update_past_key_values_sw(self, past_key_values, window_size):
649
+ past_key_values = past_key_values.to_legacy_cache()
650
+ past_key_values = [
651
+ [
652
+ k_or_v[..., -(window_size+self.use_sink):, :]
653
+ for k_or_v in seg_kv
654
+ ]
655
+ for seg_kv in past_key_values
656
+ ]
657
+ past_key_values = DynamicCache.from_legacy_cache(past_key_values)
658
+ return past_key_values
659
+
660
+ def greedy_generate_sw(self, input_ids, attention_mask, prev_attn_mask, **generate_kwargs):
661
+ self.generate_mode(True)
662
+ window_size = generate_kwargs['window_size']
663
+ max_new_tokens = generate_kwargs['max_new_tokens']
664
+ past_key_values = self.update_past_key_values_sw(generate_kwargs['past_key_values'], window_size)
665
+ eos_token_id = generate_kwargs['eos_token_id']
666
+ prev_attn_mask_2d = prev_attn_mask.clone()
667
+ attention_mask_2d = attention_mask.clone()
668
+
669
+ attention_mask = attn_mask_to_4d(attention_mask, upper=False, query_len=attention_mask.size(-1))
670
+ prev_attn_mask = attn_mask_to_4d(prev_attn_mask, upper=True, query_len=attention_mask.size(-1))
671
+ seg_kwargs = self.process_input(input_ids=input_ids, attention_mask=attention_mask, prev_attn_mask=prev_attn_mask, past_key_values=past_key_values)
672
+ seg_kwargs['inputs_embeds'] = seg_kwargs['inputs_embeds'][..., :-self.num_mem_tokens, :]
673
+ seg_kwargs['attention_mask'] = seg_kwargs['attention_mask'][..., :-self.num_mem_tokens, :-self.num_mem_tokens]
674
+ outputs = self.model(**seg_kwargs, use_cache=True)
675
+
676
+ next_token_logits = outputs.logits[:, -1, :]
677
+
678
+ past_key_values = outputs.past_key_values
679
+ past_key_values = self.update_past_key_values_sw(past_key_values, window_size)
680
+
681
+ generated_ids = None
682
+ sw_attention_mask = torch.cat([prev_attn_mask_2d, torch.ones(attention_mask_2d.size(0), 1).to(prev_attn_mask_2d.device), attention_mask_2d], dim=-1)
683
+
684
+ for i in range(max_new_tokens):
685
+ # print(next_token_logits[..., :5])
686
+ next_token_id = torch.argmax(next_token_logits, dim=-1).unsqueeze(-1)
687
+
688
+ if generated_ids is not None:
689
+ generated_ids = torch.cat([generated_ids, next_token_id], dim=-1)
690
+ else:
691
+ generated_ids = next_token_id
692
+ next_input = next_token_id
693
+
694
+ sw_attention_mask = torch.cat([sw_attention_mask, torch.ones_like(next_token_id).to(sw_attention_mask.device)], dim=-1)[..., -window_size-1-self.use_sink:]
695
+ with torch.no_grad():
696
+ outputs = self.model(
697
+ input_ids=next_input,
698
+ attention_mask=sw_attention_mask,
699
+ past_key_values=past_key_values,
700
+ use_cache=True,
701
+ cache_position=torch.full((1,), window_size + i + input_ids.size(-1) + self.use_sink).to(input_ids.device)
702
+ )
703
+ past_key_values = self.update_past_key_values_sw(outputs.past_key_values, window_size)
704
+ next_token_logits = outputs.logits[:, -1, :]
705
+ #print(outputs.logits.shape)
706
+ if (next_token_id[:, 0] == eos_token_id).all():
707
+ break
708
+ self.generate_mode(False)
709
+ return generated_ids
710
+
711
+ def greedy_generate_sw_shift(self, input_ids, attention_mask, prev_attn_mask, **generate_kwargs):
712
+ self.generate_mode(True)
713
+ print("Enabled generate mode, shifted gen")
714
+ window_size = generate_kwargs['window_size']
715
+ max_new_tokens = generate_kwargs['max_new_tokens']
716
+ # past_key_values = self.update_past_key_values_sw(generate_kwargs['past_key_values'], window_size)
717
+ eos_token_id = generate_kwargs['eos_token_id']
718
+
719
+ generated_ids = input_ids[..., :-1]
720
+ initial_length = input_ids.shape[-1]
721
+ past_key_values = self.update_past_key_values_sw(generate_kwargs['past_key_values'], window_size-initial_length)
722
+
723
+ # sw_attention_mask = torch.cat([prev_attn_mask[..., -window_size:], attention_mask[..., :-1]], dim=-1)[..., -window_size-initial_length:]
724
+ sw_attention_mask = torch.cat([prev_attn_mask[..., -window_size:], attention_mask[..., :-1]], dim=-1)[..., -window_size:]
725
+ #print(sw_attention_mask.shape)
726
+
727
+ for i in range(input_ids.size(-1)-1, input_ids.size(-1) + max_new_tokens):
728
+
729
+ if i < input_ids.size(-1):
730
+ next_token_id = input_ids[..., i:i+1]
731
+ else:
732
+ next_token_id = torch.argmax(next_token_logits, dim=-1).unsqueeze(-1)
733
+
734
+ if generated_ids is not None and i >= input_ids.size(-1)-1:
735
+ generated_ids = torch.cat([generated_ids, next_token_id], dim=-1)
736
+ else:
737
+ generated_ids = next_token_id
738
+ # TODO: think how to fix this trunc to initial len
739
+ # next_input = generated_ids[..., -initial_length:]
740
+ next_input = generated_ids[..., -window_size:]
741
+ #if next_input.shape[-1] > window_size:
742
+ # next_input = next_input[..., -window_size:]
743
+ # TODO: check attn mask - maybe it's partially inf, and partially non inf - no, all mask is ones
744
+ if i < input_ids.size(-1):
745
+ # sw_attention_mask = torch.cat([sw_attention_mask, attention_mask[..., i:i+1]], dim=-1)[..., -window_size-initial_length:]
746
+ sw_attention_mask = torch.cat([sw_attention_mask, attention_mask[..., i:i+1]], dim=-1)[..., -window_size:]
747
+ else:
748
+ # sw_attention_mask = torch.cat([sw_attention_mask, torch.ones_like(next_token_id)], dim=-1)[..., -window_size-initial_length:]
749
+ sw_attention_mask = torch.cat([sw_attention_mask, torch.ones_like(next_token_id)], dim=-1)[..., -window_size:]
750
+ #print(sw_attention_mask)
751
+ #print(input_ids.shape, next_input.shape, sw_attention_mask.shape, past_key_values)
752
+ #print(past_key_values[-1][0].shape)
753
+ with torch.no_grad():
754
+ outputs = self.model(
755
+ input_ids=next_input,
756
+ attention_mask=sw_attention_mask,
757
+ past_key_values=past_key_values,
758
+ use_cache=True
759
+ )
760
+ # past_key_values = self.update_past_key_values_sw(outputs.past_key_values, window_size)
761
+ past_key_values = self.update_past_key_values_sw(outputs.past_key_values, window_size-initial_length)
762
+ # TODO: check logits selection
763
+ #print(outputs.logits.shape)
764
+ next_token_logits = outputs.logits[:, -1, :]
765
+ if (next_token_id[:, 0] == eos_token_id).all():
766
+ break
767
+ #print(smth)
768
+ #print(input_ids)
769
+ #print(generated_ids)
770
+ self.generate_mode(False)
771
+ return generated_ids[..., initial_length:]
772
+
773
+ def greedy_generate_sw_my(self, input_ids, attention_mask, **generate_kwargs):
774
+ window_size = generate_kwargs['window_size']
775
+ max_new_tokens = generate_kwargs['max_new_tokens']
776
+ past_key_values = self.update_past_key_values_sw(generate_kwargs['past_key_values'], window_size)
777
+ eos_token_id = generate_kwargs['eos_token_id']
778
+
779
+ generated_ids = input_ids[..., :-1] #None
780
+ attention_mask = attention_mask[..., :-1]
781
+
782
+ #for i in range(input_ids.size(-1) + max_new_tokens):
783
+ print(input_ids)
784
+ for i in range(input_ids.size(-1)-1, input_ids.size(-1) + max_new_tokens):
785
+
786
+ if i < input_ids.size(-1):
787
+ next_token_id = input_ids[..., i:i+1]
788
+ else:
789
+ next_token_id = torch.argmax(next_token_logits, dim=-1).unsqueeze(-1)
790
+
791
+ if generated_ids is not None and i >= input_ids.size(-1)-1:
792
+ generated_ids = torch.cat([generated_ids, next_token_id], dim=-1)
793
+ else:
794
+ generated_ids = next_token_id
795
+ next_input = generated_ids
796
+ print(next_input)
797
+ attention_mask = torch.cat([attention_mask, torch.ones_like(next_token_id)], dim=-1)
798
+ with torch.no_grad():
799
+ print(input_ids.shape, next_input.shape, attention_mask.shape, past_key_values)
800
+ outputs = self.model(
801
+ input_ids=next_input,
802
+ attention_mask=attention_mask,
803
+ past_key_values=past_key_values,
804
+ use_cache=True
805
+ )
806
+ past_key_values = self.update_past_key_values_sw(outputs.past_key_values, window_size)
807
+ next_token_logits = outputs.logits[:, -1, :]
808
+ if (next_token_id[:, 0] == eos_token_id).all():
809
+ break
810
+ return generated_ids
811
+
812
+
813
+ class AssociativeRecurrentWrapper(torch.nn.Module):
814
+ def __init__(self, memory_cell, **rmt_kwargs):
815
+ super().__init__()
816
+
817
+ self.memory_cell = memory_cell
818
+ self.rmt_config = rmt_kwargs
819
+
820
+ def gradient_checkpointing_enable(self, *args, **kwargs):
821
+ self.memory_cell.model.gradient_checkpointing_enable(*args, **kwargs)
822
+
823
+ def process_segment(self, segment_kwargs, next_seg_len=None):
824
+ sliding_window = self.rmt_config['sliding_window'] if 'sliding_window' in self.rmt_config else False
825
+ attend_to_previous_input = self.rmt_config['attend_to_previous_input'] if 'attend_to_previous_input' in self.rmt_config else False
826
+ attn_mask = segment_kwargs['attention_mask']
827
+ seg_len = segment_kwargs['input_ids'].size(-1)
828
+
829
+ segment_kwargs['use_cache'] = sliding_window
830
+ if segment_kwargs.get('past_key_values') is None:
831
+ segment_kwargs['past_key_values'] = None
832
+ if segment_kwargs.get('prev_attn_mask') is None:
833
+ segment_kwargs['prev_attn_mask'] = None
834
+ segment_kwargs['zero_mem'] = False
835
+ if sliding_window or attend_to_previous_input:
836
+ segment_kwargs['attention_mask'] = attn_mask_to_4d(attn_mask, upper=False, query_len=seg_len)
837
+
838
+
839
+ num_mem_tokens = self.memory_cell.num_mem_tokens
840
+ cell_out = self.memory_cell(**segment_kwargs)
841
+ state = cell_out.get('state')
842
+ if (sliding_window or attend_to_previous_input) and next_seg_len is not None:
843
+ prev_attn_mask = attn_mask_to_4d(attn_mask, upper=True, query_len=next_seg_len)
844
+ else:
845
+ prev_attn_mask = None
846
+ if sliding_window:
847
+ past_key_values = [
848
+ [
849
+ k_or_v[..., -(num_mem_tokens+seg_len):k_or_v.size(-2)-num_mem_tokens, :].detach()
850
+ for k_or_v in seg_kv
851
+ ]
852
+ for seg_kv in cell_out['past_key_values']
853
+ ]
854
+ if not isinstance(cell_out['past_key_values'], tuple) and not isinstance(cell_out['past_key_values'], list):
855
+ past_key_values = cell_out['past_key_values'].from_legacy_cache(past_key_values)
856
+ else:
857
+ past_key_values = DynamicCache.from_legacy_cache(past_key_values)
858
+ else:
859
+ past_key_values = None
860
+ next_segment_kwargs = dict()
861
+ next_segment_kwargs['use_cache'] = sliding_window
862
+ next_segment_kwargs['past_key_values'] = past_key_values
863
+ next_segment_kwargs['prev_attn_mask'] = prev_attn_mask
864
+ next_segment_kwargs['zero_mem'] = False
865
+ if state is not None:
866
+ next_segment_kwargs['state'] = state
867
+ return cell_out, next_segment_kwargs
868
+
869
+ def process_last_segment(self, segment_kwargs, next_seg_len=None):
870
+ sliding_window = self.rmt_config['sliding_window'] if 'sliding_window' in self.rmt_config else False
871
+ attend_to_previous_input = self.rmt_config['attend_to_previous_input'] if 'attend_to_previous_input' in self.rmt_config else False
872
+ attn_mask = segment_kwargs['attention_mask']
873
+ seg_len = segment_kwargs['input_ids'].size(-1)
874
+
875
+ segment_kwargs['use_cache'] = sliding_window
876
+ if segment_kwargs.get('past_key_values') is None:
877
+ segment_kwargs['past_key_values'] = None
878
+ if segment_kwargs.get('prev_attn_mask') is None:
879
+ segment_kwargs['prev_attn_mask'] = None
880
+ segment_kwargs['zero_mem'] = False
881
+ if sliding_window or attend_to_previous_input:
882
+ segment_kwargs['attention_mask'] = self.attn_mask_to_4d(attn_mask, upper=False, query_len=seg_len)
883
+
884
+ if segment_kwargs.get('prev_attn_mask') is not None:
885
+ print("Prev attn mask start", segment_kwargs['prev_attn_mask'].shape)
886
+ num_mem_tokens = self.memory_cell.num_mem_tokens
887
+ cell_out = self.memory_cell(**segment_kwargs)
888
+ state = cell_out.get('state')
889
+ # simply keep prev attn mask
890
+ #if (sliding_window or attend_to_previous_input) and next_seg_len is not None:
891
+ # prev_attn_mask = self.attn_mask_to_4d(attn_mask, upper=True, query_len=next_seg_len)
892
+ #else:
893
+ # prev_attn_mask = None
894
+ if sliding_window:
895
+ print("Past key vals start", cell_out['past_key_values'][0][0].shape)
896
+ past_key_values = [
897
+ [
898
+ k_or_v[..., 1:k_or_v.size(-2)-num_mem_tokens-seg_len+1, :].detach()
899
+ for k_or_v in seg_kv
900
+ ]
901
+ for seg_kv in cell_out['past_key_values']
902
+ ]
903
+ if not isinstance(cell_out['past_key_values'], tuple) and not isinstance(cell_out['past_key_values'], list):
904
+ past_key_values = cell_out['past_key_values'].from_legacy_cache(past_key_values)
905
+ else:
906
+ past_key_values = DynamicCache.from_legacy_cache(past_key_values)
907
+ else:
908
+ past_key_values = None
909
+ next_segment_kwargs = dict()
910
+ next_segment_kwargs['use_cache'] = sliding_window
911
+ next_segment_kwargs['past_key_values'] = past_key_values
912
+ next_segment_kwargs['prev_attn_mask'] = segment_kwargs['prev_attn_mask']
913
+ next_segment_kwargs['zero_mem'] = False
914
+ if state is not None:
915
+ next_segment_kwargs['state'] = state
916
+ return cell_out, next_segment_kwargs
917
+
918
+ def forward(self,
919
+ input_ids,
920
+ labels=None,
921
+ labels_mask=None,
922
+ inputs_embeds=None,
923
+ attention_mask=None,
924
+ output_attentions=None,
925
+ output_hidden_states=None,
926
+ input_segmented=False,
927
+ output_only_last_segment=False,
928
+ ):
929
+ if input_segmented:
930
+ n_segs = input_ids.shape[1] if not (input_ids is None) else inputs_embeds.shape[1]
931
+ segmented = [dict(
932
+ input_ids=input_ids[:, i] if not (input_ids is None) else None,
933
+ inputs_embeds=inputs_embeds[:, i] if not (inputs_embeds is None) else None,
934
+ attention_mask=attention_mask[:, i],
935
+ labels=labels[:, i] if not (labels is None) else None,
936
+ labels_mask=labels_mask[:, i] if not (labels_mask is None) else None,
937
+ ) for i in range(n_segs)]
938
+ labels = torch.cat([labels[:, i] for i in range(n_segs)], dim=1)
939
+ if labels_mask is not None:
940
+ labels_mask = torch.cat([labels_mask[:, i] for i in range(n_segs)], dim=1)
941
+ else:
942
+ segmented = self.segment(input_ids=input_ids, inputs_embeds=inputs_embeds, attention_mask=attention_mask, labels=labels, labels_mask=labels_mask)
943
+
944
+ cell_outputs = []
945
+ self.memory_cell.zero_mem()
946
+ next_seg_kwargs = dict()
947
+ for seg_num, segment in enumerate(segmented):
948
+ if seg_num != len(segmented) - 1:
949
+ next_seg_len = segmented[seg_num + 1]['input_ids'].size(-1)
950
+ else:
951
+ next_seg_len = None
952
+ cell_out, next_seg_kwargs = self.process_segment(dict(**segment, **next_seg_kwargs), next_seg_len=next_seg_len)
953
+ if (not output_only_last_segment) or (seg_num == len(segmented) - 1):
954
+ cell_outputs.append(cell_out)
955
+
956
+ out = self.process_outputs(cell_outputs, labels=labels,
957
+ labels_mask=labels_mask,
958
+ output_attentions=output_attentions,
959
+ output_hidden_states=output_hidden_states)
960
+ return out
961
+
962
+ def segment(self, **kwargs):
963
+ segments = []
964
+ for k, tensor in kwargs.items():
965
+ if tensor is not None:
966
+ k_segments = self.split_tensor(tensor)
967
+ for s, k_seg in enumerate(k_segments):
968
+ if s < len(segments):
969
+ segments[s][k] = k_seg
970
+ else:
971
+ segments.append({k: k_seg})
972
+
973
+ return segments
974
+
975
+ def split_tensor(self, tensor):
976
+ align = self.rmt_config.get('segment_alignment')
977
+ segment_size = self.rmt_config.get('segment_size')
978
+ if align in {'left', None}:
979
+ split_inds = list(range(0, tensor.shape[1], segment_size)) + [tensor.shape[1]]
980
+ segments = [tensor[:, start:end] for (start, end) in zip(split_inds, split_inds[1:])]
981
+ elif align in {'right', None}:
982
+ split_inds = (list(range(tensor.shape[1], 0, -segment_size)) + [0])[::-1]
983
+ segments = [tensor[:, start:end] for (start, end) in zip(split_inds, split_inds[1:])]
984
+ elif align == 'center':
985
+ n_seg = math.ceil(tensor.shape[1] / segment_size)
986
+ segments = torch.chunk(tensor, n_seg, dim=1)
987
+ else:
988
+ raise NotImplementedError
989
+ return segments
990
+
991
+ def process_outputs(self, cell_outputs, **kwargs):
992
+ out = CausalLMOutputWithCrossAttentions()
993
+ full_logits = torch.cat([o.logits for o in cell_outputs], dim=1)
994
+
995
+ labels = kwargs.get('labels')
996
+ if labels is not None:
997
+ labels = labels[:, -full_logits.size(1):]
998
+ shift_labels = labels[..., 1:].contiguous()
999
+ shift_logits = full_logits[..., :-1, :].contiguous()
1000
+ flat_labels = shift_labels.view(-1)
1001
+ flat_logits = shift_logits.view(-1, shift_logits.size(-1))
1002
+
1003
+ loss_fct = CrossEntropyLoss()
1004
+ labels_mask = kwargs.get('labels_mask')
1005
+ if labels_mask is not None:
1006
+ labels_mask = labels_mask[:, -full_logits.size(1):]
1007
+ shift_mask = labels_mask[..., :-1].contiguous()
1008
+
1009
+ flat_labels = flat_labels[shift_mask.view(-1)]
1010
+ flat_logits = flat_logits[shift_mask.view(-1)]
1011
+
1012
+ out['loss'] = loss_fct(flat_logits, flat_labels)
1013
+ else:
1014
+ out['loss'] = 0
1015
+ if (('HF_Trainer' not in os.environ) or not os.environ['HF_Trainer']) and self.rmt_config.get("return_all_logits", False):
1016
+ out['ce_loss'] = out['loss']
1017
+
1018
+ out['logits'] = full_logits
1019
+ segment_keys = ['loss', 'logits']
1020
+ if kwargs.get('output_attentions'):
1021
+ segment_keys.append('attentions')
1022
+ if kwargs.get('output_hidden_states'):
1023
+ full_hidden_states = tuple([torch.cat(layer_hs, dim=1) for layer_hs in zip(*[o.hidden_states for o in cell_outputs])])
1024
+ segment_keys.append('hidden_states')
1025
+ out['hidden_states'] = full_hidden_states
1026
+ if (('HF_Trainer' not in os.environ) or not os.environ['HF_Trainer']) and self.rmt_config.get("return_all_logits", False):
1027
+ for seg_num, o in enumerate(cell_outputs):
1028
+ for key, value in o.items():
1029
+ if any([sk in key for sk in segment_keys]):
1030
+ out[f'{key}_{seg_num}'] = value
1031
+
1032
+ remainders = []
1033
+ n_updates = []
1034
+ act_on = self.rmt_config['act_on'] if 'act_on' in self.rmt_config else False
1035
+ if act_on:
1036
+
1037
+ for layer in self.memory_cell.layers:
1038
+ remainders.append(layer.remainders / layer.segments_passed)
1039
+ n_updates.append(layer.n_updates / layer.segments_passed)
1040
+ remainders = torch.mean(torch.stack(remainders, dim=0))
1041
+ n_updates = torch.mean(torch.stack(n_updates, dim=0))
1042
+ out['n_updates'] = n_updates.detach().cpu()
1043
+ out['remainders'] = remainders.detach().cpu()
1044
+ time_penalty = self.rmt_config['time_penalty']
1045
+ out['loss'] = out['loss'] + time_penalty * remainders
1046
+
1047
+ return out
1048
+
1049
+ def manage_gradients(self, memory_state, seg_num):
1050
+ k2, max_n_segments = self.rmt_config.get('k2'), self.rmt_config.get('max_n_segments')
1051
+ if seg_num == 0 \
1052
+ or k2 in {-1, None} \
1053
+ or seg_num + k2 > max_n_segments:
1054
+ return True
1055
+
1056
+ memory_state = memory_state.detach()
1057
+ return False
1058
+
1059
+ def generate(self, input_ids, attention_mask, **generate_kwargs):
1060
+ #print(input_ids.shape, attention_mask.shape)
1061
+ self.memory_cell.zero_mem()
1062
+ segmented = self.segment(input_ids=input_ids, attention_mask=attention_mask)
1063
+ next_seg_kwargs = dict()
1064
+ for seg_num, segment in enumerate(segmented[:-1]):
1065
+ next_seg_len = segmented[seg_num + 1]['input_ids'].size(-1)
1066
+ _, next_seg_kwargs = self.process_segment(dict(**segment, **next_seg_kwargs), next_seg_len=next_seg_len)
1067
+
1068
+ final_segment = segmented[-1]
1069
+ assert next_seg_kwargs.get('past_key_values') is None or isinstance(next_seg_kwargs.get('past_key_values'), Cache), "Sliding Window generation is not implemented for legacy cache"
1070
+ if next_seg_kwargs.get('past_key_values') is not None:
1071
+ """
1072
+ prev_attn_mask = segmented[-2]['attention_mask']
1073
+ legacy_cache = next_seg_kwargs['past_key_values']
1074
+ seg_len = segmented[-2]['input_ids'].size(-1)
1075
+ #cache = DynamicCache().from_legacy_cache(legacy_cache)
1076
+ generate_kwargs['past_key_values'] = legacy_cache
1077
+ generate_kwargs['window_size'] = seg_len
1078
+ #final_segment['prev_attn_mask'] = self.attn_mask_to_4d(prev_attn_mask, upper=True, query_len=seg_len)
1079
+ #del next_seg_kwargs["prev_attn_mask"]
1080
+ print(final_segment.keys())
1081
+ print(final_segment["input_ids"].shape)
1082
+ print(next_seg_kwargs["past_key_values"][0][0].shape)
1083
+ max_tokens = generate_kwargs["max_new_tokens"]
1084
+ generations = None
1085
+ for idx in range(max_tokens):
1086
+ # TODO: shift past_kv_values on one step, and add new ids to the input
1087
+ cell_out, next_seg_kwargs = self.process_last_segment(dict(**final_segment, **next_seg_kwargs), next_seg_len=seg_len)
1088
+ print(next_seg_kwargs["past_key_values"][0][0].shape)
1089
+ print(cell_out.logits.shape)
1090
+ #print(final_segment["input_ids"])
1091
+ print(torch.argmax(cell_out.logits[:, -1, :], dim=-1))
1092
+ next_token = torch.argmax(cell_out.logits[:, -1, :], dim=-1)
1093
+ if generations is None:
1094
+ generations = next_token.unsqueeze(1)
1095
+ else:
1096
+ generations = torch.cat([generations, next_token.unsqueeze(1)], dim=1)
1097
+ final_segment["input_ids"] = torch.cat([final_segment["input_ids"], next_token.unsqueeze(1)], dim=1)[..., 1:]
1098
+ print(final_segment["input_ids"].shape)
1099
+ if next_token == generate_kwargs["eos_token_id"]:
1100
+ break
1101
+ #out = self.memory_cell.greedy_generate_sw(**final_segment, **generate_kwargs)
1102
+ return generations
1103
+ """
1104
+ prev_attn_mask = segmented[-2]['attention_mask']
1105
+ legacy_cache = next_seg_kwargs['past_key_values'].to_legacy_cache()
1106
+ seg_len = segmented[-2]['input_ids'].size(-1)
1107
+ cache = DynamicCache().from_legacy_cache(legacy_cache)
1108
+ generate_kwargs['past_key_values'] = cache
1109
+ generate_kwargs['window_size'] = seg_len
1110
+ final_segment['prev_attn_mask'] = prev_attn_mask
1111
+ out = self.memory_cell.greedy_generate_sw(**final_segment, **generate_kwargs)
1112
+ return out
1113
+ else:
1114
+ out = self.memory_cell.generate(**final_segment, **generate_kwargs)
1115
+ return out
1116
+
1117
+ def generate_prom(self, input_ids, attention_mask, **generate_kwargs):
1118
+ sliding_window = self.rmt_config['sliding_window'] if 'sliding_window' in self.rmt_config else False
1119
+ self.memory_cell.zero_mem()
1120
+ segmented = self.segment(input_ids=input_ids, attention_mask=attention_mask)
1121
+
1122
+ num_mem_tokens = self.memory_cell.num_mem_tokens
1123
+ past_key_values = None
1124
+ prev_attn_mask = None
1125
+ for seg_num, segment in enumerate(segmented[:-1]):
1126
+ seg_len = segment['input_ids'].size(-1)
1127
+ segment['use_cache'] = sliding_window
1128
+ segment['past_key_values'] = past_key_values
1129
+ segment['prev_attn_mask'] = prev_attn_mask
1130
+ attn_mask = segment['attention_mask']
1131
+ if sliding_window:
1132
+ segment['attention_mask'] = self.attn_mask_to_4d(attn_mask, upper=False, query_len=seg_len)
1133
+ cell_out = self.memory_cell(**segment, output_hidden_states=True, zero_mem=False)
1134
+ if sliding_window and seg_num + 1 != len(segmented):
1135
+ next_seg_len = segmented[seg_num+1]['input_ids'].size(-1)
1136
+ prev_attn_mask = self.attn_mask_to_4d(attn_mask, upper=True, query_len=next_seg_len)
1137
+ if sliding_window:
1138
+ past_key_values = [
1139
+ [
1140
+ k_or_v[..., -(num_mem_tokens+seg_len):k_or_v.size(-2)-num_mem_tokens, :].detach()
1141
+ for k_or_v in seg_kv
1142
+ ]
1143
+ for seg_kv in cell_out['past_key_values']
1144
+ ]
1145
+ if not isinstance(cell_out['past_key_values'], tuple) and not isinstance(cell_out['past_key_values'], list):
1146
+ past_key_values = cell_out['past_key_values'].from_legacy_cache(past_key_values)
1147
+ final_segment = segmented[-1]
1148
+ #seg_len = final_segment['input_ids'].size(-1)
1149
+ #final_segment['use_cache'] = sliding_window
1150
+ #final_segment['past_key_values'] = past_key_values
1151
+ #final_segment['prev_attn_mask'] = prev_attn_mask
1152
+ #attn_mask = final_segment['attention_mask']
1153
+ #if sliding_window:
1154
+ # final_segment['attention_mask'] = self.attn_mask_to_4d(attn_mask, upper=False, query_len=seg_len)
1155
+ out = self.memory_cell.generate(**final_segment, zero_mem=False, **generate_kwargs)
1156
+ self.memory_cell.zero_mem()
1157
+ return out