userkuku commited on
Commit
ecc111a
·
verified ·
1 Parent(s): cedd68d

Upload model_architecture.py

Browse files
Files changed (1) hide show
  1. model_architecture.py +616 -0
model_architecture.py ADDED
@@ -0,0 +1,616 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ LMCODE: Language Model with Memory CODE
3
+ =======================================
4
+
5
+ A memory-augmented language model with dual memory systems:
6
+ - Short-term memory: Working memory for immediate context (like Transformer KV cache)
7
+ - Long-term memory: Persistent storage for knowledge and experiences
8
+
9
+ Inspired by:
10
+ - LongMem (2023): Augmenting LLMs with Long-Term Memory
11
+ - MemoRAG (2024): Dual-system RAG with Global and Local Memory
12
+ - CAMELoT (2024): Training-free Consolidated Associative Memory
13
+ """
14
+
15
+ import torch
16
+ import torch.nn as nn
17
+ import torch.nn.functional as F
18
+ from typing import Optional, Tuple, Dict, List
19
+ import math
20
+
21
+
22
+ class ShortTermMemory(nn.Module):
23
+ """
24
+ Short-term memory module using recurrent state updates.
25
+ Acts as working memory for immediate context (similar to Transformer KV cache).
26
+ """
27
+
28
+ def __init__(self, hidden_size: int, memory_size: int = 512):
29
+ super().__init__()
30
+ self.hidden_size = hidden_size
31
+ self.memory_size = memory_size
32
+
33
+ # Memory state
34
+ self.memory = nn.Parameter(torch.zeros(1, memory_size, hidden_size))
35
+
36
+ # Memory update gate
37
+ self.update_gate = nn.Linear(hidden_size, 1)
38
+
39
+ # Memory read/write projections
40
+ self.read_proj = nn.Linear(hidden_size, hidden_size)
41
+ self.write_proj = nn.Linear(hidden_size, hidden_size)
42
+
43
+ # Initialize memory
44
+ nn.init.normal_(self.memory, mean=0, std=0.02)
45
+
46
+ def forward(self, hidden_states: torch.Tensor,
47
+ attention_mask: Optional[torch.Tensor] = None) -> Dict[str, torch.Tensor]:
48
+ """
49
+ Update and read from short-term memory.
50
+
51
+ Args:
52
+ hidden_states: Current hidden states [batch_size, seq_len, hidden_size]
53
+ attention_mask: Optional attention mask
54
+
55
+ Returns:
56
+ Dictionary with memory outputs
57
+ """
58
+ batch_size, seq_len, _ = hidden_states.shape
59
+
60
+ # Expand memory for batch
61
+ memory = self.memory.expand(batch_size, -1, -1) # [batch, memory_size, hidden]
62
+
63
+ # Compute update gates
64
+ update_scores = torch.sigmoid(self.update_gate(hidden_states)) # [batch, seq, 1]
65
+
66
+ # Read from memory
67
+ read_query = self.read_proj(hidden_states) # [batch, seq, hidden]
68
+ read_scores = torch.matmul(read_query, memory.transpose(-1, -2)) # [batch, seq, memory_size]
69
+ read_scores = read_scores / math.sqrt(self.hidden_size)
70
+
71
+ if attention_mask is not None:
72
+ read_scores = read_scores.masked_fill(~attention_mask.unsqueeze(-1).bool(), -1e9)
73
+
74
+ read_weights = F.softmax(read_scores, dim=-1) # [batch, seq, memory_size]
75
+ memory_read = torch.matmul(read_weights, memory) # [batch, seq, hidden]
76
+
77
+ # Write to memory (soft update)
78
+ write_values = self.write_proj(hidden_states) # [batch, seq, hidden]
79
+ write_weights = update_scores.expand(-1, -1, self.memory_size) # [batch, seq, memory_size]
80
+
81
+ # Update memory (detached to prevent gradient flow through time)
82
+ memory_delta = torch.matmul(write_weights.transpose(-1, -2), write_values) # [batch, memory_size, hidden]
83
+ memory = memory + 0.1 * memory_delta # Small learning rate for memory update
84
+
85
+ # Combine hidden states with memory read
86
+ output = hidden_states + 0.5 * memory_read # Residual connection
87
+
88
+ return {
89
+ 'output': output,
90
+ 'memory_read': memory_read,
91
+ 'read_weights': read_weights,
92
+ 'updated_memory': memory
93
+ }
94
+
95
+
96
+ class LongTermMemory(nn.Module):
97
+ """
98
+ Long-term memory module for persistent storage of knowledge.
99
+ Uses a key-value store with retrieval mechanism.
100
+
101
+ Inspired by:
102
+ - LongMem's adaptive residual side-network
103
+ - MemoRAG's global memory for overview
104
+ """
105
+
106
+ def __init__(self, hidden_size: int, num_slots: int = 10000):
107
+ super().__init__()
108
+ self.hidden_size = hidden_size
109
+ self.num_slots = num_slots
110
+
111
+ # Memory storage (key-value pairs)
112
+ self.memory_keys = nn.Parameter(torch.zeros(num_slots, hidden_size))
113
+ self.memory_values = nn.Parameter(torch.zeros(num_slots, hidden_size))
114
+
115
+ # Memory metadata (importance, timestamps, etc.)
116
+ self.memory_importance = nn.Parameter(torch.ones(num_slots))
117
+
118
+ # Query projection
119
+ self.query_proj = nn.Linear(hidden_size, hidden_size)
120
+ self.key_proj = nn.Linear(hidden_size, hidden_size)
121
+ self.value_proj = nn.Linear(hidden_size, hidden_size)
122
+
123
+ # Memory consolidation (for summarizing old memories)
124
+ self.consolidation_layer = nn.Sequential(
125
+ nn.Linear(hidden_size * 2, hidden_size),
126
+ nn.ReLU(),
127
+ nn.Linear(hidden_size, hidden_size)
128
+ )
129
+
130
+ # Initialize
131
+ nn.init.normal_(self.memory_keys, mean=0, std=0.02)
132
+ nn.init.normal_(self.memory_values, mean=0, std=0.02)
133
+ nn.init.constant_(self.memory_importance, 1.0)
134
+
135
+ def forward(self, hidden_states: torch.Tensor,
136
+ store_new: bool = False,
137
+ new_content: Optional[torch.Tensor] = None) -> Dict[str, torch.Tensor]:
138
+ """
139
+ Retrieve from or store to long-term memory.
140
+
141
+ Args:
142
+ hidden_states: Current hidden states [batch_size, seq_len, hidden_size]
143
+ store_new: Whether to store new content
144
+ new_content: New content to store [batch_size, content_len, hidden_size]
145
+
146
+ Returns:
147
+ Dictionary with memory outputs
148
+ """
149
+ batch_size, seq_len, _ = hidden_states.shape
150
+
151
+ # Project queries
152
+ queries = self.query_proj(hidden_states) # [batch, seq, hidden]
153
+
154
+ # Project memory keys
155
+ keys = self.key_proj(self.memory_keys) # [num_slots, hidden]
156
+
157
+ # Compute retrieval scores
158
+ scores = torch.matmul(queries, keys.transpose(-1, -2)) # [batch, seq, num_slots]
159
+ scores = scores / math.sqrt(self.hidden_size)
160
+
161
+ # Apply importance weighting
162
+ importance_weights = torch.sigmoid(self.memory_importance) # [num_slots]
163
+ scores = scores * importance_weights.unsqueeze(0).unsqueeze(0) # [batch, seq, num_slots]
164
+
165
+ # Get top-k memories
166
+ top_k = min(10, self.num_slots)
167
+ top_scores, top_indices = torch.topk(scores, k=top_k, dim=-1) # [batch, seq, top_k]
168
+
169
+ # Softmax over top-k
170
+ retrieval_weights = F.softmax(top_scores, dim=-1) # [batch, seq, top_k]
171
+
172
+ # Gather retrieved memories
173
+ retrieved_values = self.memory_values[top_indices] # [batch, seq, top_k, hidden]
174
+
175
+ # Weighted sum
176
+ memory_output = torch.sum(retrieved_values * retrieval_weights.unsqueeze(-1), dim=-2) # [batch, seq, hidden]
177
+
178
+ # Project memory output
179
+ memory_output = self.value_proj(memory_output)
180
+
181
+ # Store new memories if requested
182
+ if store_new and new_content is not None:
183
+ self._store_new_memories(new_content)
184
+
185
+ return {
186
+ 'output': memory_output,
187
+ 'retrieval_weights': retrieval_weights,
188
+ 'top_indices': top_indices,
189
+ 'retrieval_scores': top_scores
190
+ }
191
+
192
+ def _store_new_memories(self, content: torch.Tensor):
193
+ """
194
+ Store new content in long-term memory.
195
+ Uses a simple FIFO with consolidation strategy.
196
+
197
+ Args:
198
+ content: Content to store [batch_size, content_len, hidden_size]
199
+ """
200
+ with torch.no_grad():
201
+ # Average content across batch and sequence
202
+ content_repr = content.mean(dim=(0, 1)) # [hidden_size]
203
+
204
+ # Find least important slot
205
+ _, least_important_idx = torch.min(self.memory_importance, dim=0)
206
+
207
+ # Update memory (with small learning rate)
208
+ lr = 0.01
209
+ self.memory_keys.data[least_important_idx] = \
210
+ (1 - lr) * self.memory_keys.data[least_important_idx] + lr * content_repr
211
+ self.memory_values.data[least_important_idx] = \
212
+ (1 - lr) * self.memory_values.data[least_important_idx] + lr * content_repr
213
+
214
+ # Reset importance (new memories start neutral)
215
+ self.memory_importance.data[least_important_idx] = 1.0
216
+
217
+ def consolidate_memories(self, threshold: float = 0.1):
218
+ """
219
+ Consolidate similar memories to prevent redundancy.
220
+
221
+ Args:
222
+ threshold: Similarity threshold for consolidation
223
+ """
224
+ with torch.no_grad():
225
+ # Compute pairwise similarity
226
+ keys_norm = F.normalize(self.memory_keys, dim=-1)
227
+ similarity = torch.matmul(keys_norm, keys_norm.transpose(-1, -2))
228
+
229
+ # Find highly similar pairs
230
+ mask = torch.triu(similarity > threshold, diagonal=1)
231
+
232
+ if mask.any():
233
+ # Consolidate: merge similar memories
234
+ indices = torch.where(mask)
235
+
236
+ for i, j in zip(indices[0], indices[1]):
237
+ # Merge j into i (weighted by importance)
238
+ imp_i = torch.sigmoid(self.memory_importance[i])
239
+ imp_j = torch.sigmoid(self.memory_importance[j])
240
+ total_imp = imp_i + imp_j
241
+
242
+ if total_imp > 0:
243
+ self.memory_keys.data[i] = \
244
+ (imp_i * self.memory_keys.data[i] + imp_j * self.memory_keys.data[j]) / total_imp
245
+ self.memory_values.data[i] = \
246
+ (imp_i * self.memory_values.data[i] + imp_j * self.memory_values.data[j]) / total_imp
247
+ self.memory_importance.data[i] = torch.log(torch.exp(self.memory_importance.data[i]) +
248
+ torch.exp(self.memory_importance.data[j]))
249
+
250
+ # Zero out merged memory
251
+ self.memory_keys.data[j] = 0
252
+ self.memory_values.data[j] = 0
253
+ self.memory_importance.data[j] = -10 # Mark as inactive
254
+
255
+
256
+ class MemoryAugmentedLayer(nn.Module):
257
+ """
258
+ Transformer layer augmented with both short-term and long-term memory.
259
+ """
260
+
261
+ def __init__(self, hidden_size: int, num_heads: int,
262
+ short_term_memory_size: int = 512,
263
+ long_term_memory_slots: int = 10000):
264
+ super().__init__()
265
+ self.hidden_size = hidden_size
266
+ self.num_heads = num_heads
267
+
268
+ # Self-attention
269
+ self.attention = nn.MultiheadAttention(hidden_size, num_heads, batch_first=True)
270
+ self.attention_norm = nn.LayerNorm(hidden_size)
271
+
272
+ # Feed-forward
273
+ self.ffn = nn.Sequential(
274
+ nn.Linear(hidden_size, hidden_size * 4),
275
+ nn.GELU(),
276
+ nn.Linear(hidden_size * 4, hidden_size)
277
+ )
278
+ self.ffn_norm = nn.LayerNorm(hidden_size)
279
+
280
+ # Memory modules
281
+ self.short_term_memory = ShortTermMemory(hidden_size, short_term_memory_size)
282
+ self.long_term_memory = LongTermMemory(hidden_size, long_term_memory_slots)
283
+
284
+ # Memory gating
285
+ self.memory_gate = nn.Linear(hidden_size * 2, hidden_size)
286
+
287
+ def forward(self, hidden_states: torch.Tensor,
288
+ attention_mask: Optional[torch.Tensor] = None,
289
+ use_long_term_memory: bool = True,
290
+ store_long_term: bool = False) -> Dict[str, torch.Tensor]:
291
+ """
292
+ Forward pass through memory-augmented layer.
293
+
294
+ Args:
295
+ hidden_states: Input hidden states [batch_size, seq_len, hidden_size]
296
+ attention_mask: Optional attention mask
297
+ use_long_term_memory: Whether to use long-term memory
298
+ store_long_term: Whether to store in long-term memory
299
+
300
+ Returns:
301
+ Dictionary with layer outputs
302
+ """
303
+ residual = hidden_states
304
+
305
+ # Self-attention
306
+ attn_output, attn_weights = self.attention(
307
+ self.attention_norm(hidden_states),
308
+ self.attention_norm(hidden_states),
309
+ self.attention_norm(hidden_states),
310
+ key_padding_mask=attention_mask
311
+ )
312
+ hidden_states = residual + attn_output
313
+
314
+ # Short-term memory
315
+ memory_output = self.short_term_memory(hidden_states, attention_mask)
316
+ hidden_states = memory_output['output']
317
+
318
+ # Long-term memory (optional)
319
+ long_term_output = None
320
+ if use_long_term_memory:
321
+ long_term_output = self.long_term_memory(
322
+ hidden_states,
323
+ store_new=store_long_term,
324
+ new_content=residual # Store original input
325
+ )
326
+ # Gate between attention and memory
327
+ memory_gate_input = torch.cat([hidden_states, long_term_output['output']], dim=-1)
328
+ gate = torch.sigmoid(self.memory_gate(memory_gate_input))
329
+ hidden_states = gate * hidden_states + (1 - gate) * long_term_output['output']
330
+
331
+ # Feed-forward
332
+ ffn_output = self.ffn(self.ffn_norm(hidden_states))
333
+ hidden_states = hidden_states + ffn_output
334
+
335
+ return {
336
+ 'hidden_states': hidden_states,
337
+ 'attention_weights': attn_weights,
338
+ 'short_term_output': memory_output,
339
+ 'long_term_output': long_term_output
340
+ }
341
+
342
+
343
+ class LMCODE(nn.Module):
344
+ """
345
+ Language Model with Memory CODE (LMCODE).
346
+
347
+ A transformer-based language model augmented with:
348
+ 1. Short-term memory: Working memory for immediate context
349
+ 2. Long-term memory: Persistent storage for knowledge and experiences
350
+
351
+ Architecture:
352
+ - Embedding layer
353
+ - Multiple memory-augmented transformer layers
354
+ - Language model head
355
+
356
+ Memory Flow:
357
+ - Short-term: Updated every forward pass, decays over time
358
+ - Long-term: Consolidated periodically, stores important experiences
359
+ """
360
+
361
+ def __init__(self, config: Optional['LMCODEConfig'] = None):
362
+ super().__init__()
363
+
364
+ if config is None:
365
+ config = LMCODEConfig()
366
+
367
+ self.config = config
368
+
369
+ # Embeddings
370
+ self.token_embeddings = nn.Embedding(config.vocab_size, config.hidden_size)
371
+ self.position_embeddings = nn.Embedding(2048, config.hidden_size)
372
+
373
+ # Memory-augmented layers
374
+ self.layers = nn.ModuleList([
375
+ MemoryAugmentedLayer(
376
+ hidden_size=config.hidden_size,
377
+ num_heads=config.num_heads,
378
+ short_term_memory_size=config.short_term_memory_size,
379
+ long_term_memory_slots=config.long_term_memory_slots
380
+ )
381
+ for _ in range(config.num_layers)
382
+ ])
383
+
384
+ # Output normalization
385
+ self.final_norm = nn.LayerNorm(config.hidden_size)
386
+
387
+ # Language model head
388
+ self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False)
389
+
390
+ # Initialize weights
391
+ self.apply(self._init_weights)
392
+
393
+ def _init_weights(self, module):
394
+ """Initialize model weights."""
395
+ if isinstance(module, nn.Linear):
396
+ nn.init.normal_(module.weight, mean=0, std=0.02)
397
+ if module.bias is not None:
398
+ nn.init.zeros_(module.bias)
399
+ elif isinstance(module, nn.Embedding):
400
+ nn.init.normal_(module.weight, mean=0, std=0.02)
401
+ elif isinstance(module, nn.LayerNorm):
402
+ nn.init.zeros_(module.bias)
403
+ nn.init.ones_(module.weight)
404
+
405
+ def forward(self, input_ids: torch.Tensor,
406
+ attention_mask: Optional[torch.Tensor] = None,
407
+ use_long_term_memory: bool = True,
408
+ store_long_term: bool = False,
409
+ labels: Optional[torch.Tensor] = None) -> Dict[str, torch.Tensor]:
410
+ """
411
+ Forward pass through LMCODE model.
412
+
413
+ Args:
414
+ input_ids: Input token IDs [batch_size, seq_len]
415
+ attention_mask: Optional attention mask [batch_size, seq_len]
416
+ use_long_term_memory: Whether to use long-term memory
417
+ store_long_term: Whether to store in long-term memory
418
+ labels: Optional labels for computing loss
419
+
420
+ Returns:
421
+ Dictionary with model outputs
422
+ """
423
+ batch_size, seq_len = input_ids.shape
424
+
425
+ # Create position IDs
426
+ position_ids = torch.arange(seq_len, device=input_ids.device).unsqueeze(0).expand(batch_size, -1)
427
+
428
+ # Embeddings
429
+ token_embeds = self.token_embeddings(input_ids)
430
+ position_embeds = self.position_embeddings(position_ids)
431
+ hidden_states = token_embeds + position_embeds
432
+
433
+ # Apply dropout
434
+ hidden_states = F.dropout(hidden_states, p=0.1, training=self.training)
435
+
436
+ # Track intermediate outputs
437
+ all_hidden_states = []
438
+ all_attention_weights = []
439
+ all_short_term_outputs = []
440
+ all_long_term_outputs = []
441
+
442
+ # Pass through layers
443
+ for layer in self.layers:
444
+ layer_output = layer(
445
+ hidden_states,
446
+ attention_mask=attention_mask,
447
+ use_long_term_memory=use_long_term_memory,
448
+ store_long_term=store_long_term
449
+ )
450
+
451
+ hidden_states = layer_output['hidden_states']
452
+ all_hidden_states.append(hidden_states)
453
+ all_attention_weights.append(layer_output['attention_weights'])
454
+ all_short_term_outputs.append(layer_output['short_term_output'])
455
+ all_long_term_outputs.append(layer_output['long_term_output'])
456
+
457
+ # Final normalization
458
+ hidden_states = self.final_norm(hidden_states)
459
+
460
+ # Language model logits
461
+ logits = self.lm_head(hidden_states)
462
+
463
+ # Compute loss if labels provided
464
+ loss = None
465
+ if labels is not None:
466
+ loss = F.cross_entropy(
467
+ logits.view(-1, self.config.vocab_size),
468
+ labels.view(-1),
469
+ ignore_index=-100
470
+ )
471
+
472
+ return {
473
+ 'logits': logits,
474
+ 'loss': loss,
475
+ 'hidden_states': all_hidden_states,
476
+ 'attention_weights': all_attention_weights,
477
+ 'short_term_outputs': all_short_term_outputs,
478
+ 'long_term_outputs': all_long_term_outputs
479
+ }
480
+
481
+ def generate(self, input_ids: torch.Tensor,
482
+ max_length: int = 100,
483
+ temperature: float = 1.0,
484
+ top_k: int = 50,
485
+ top_p: float = 0.9,
486
+ use_long_term_memory: bool = True,
487
+ store_long_term: bool = True) -> torch.Tensor:
488
+ """
489
+ Generate text autoregressively.
490
+
491
+ Args:
492
+ input_ids: Initial token IDs [batch_size, seq_len]
493
+ max_length: Maximum generation length
494
+ temperature: Sampling temperature
495
+ top_k: Top-k sampling
496
+ top_p: Nucleus sampling
497
+ use_long_term_memory: Whether to use long-term memory during generation
498
+ store_long_term: Whether to store generated text in long-term memory
499
+
500
+ Returns:
501
+ Generated token IDs [batch_size, total_length]
502
+ """
503
+ batch_size = input_ids.shape[0]
504
+
505
+ for _ in range(max_length):
506
+ # Forward pass
507
+ with torch.no_grad():
508
+ outputs = self(
509
+ input_ids,
510
+ use_long_term_memory=use_long_term_memory,
511
+ store_long_term=store_long_term
512
+ )
513
+
514
+ logits = outputs['logits'][:, -1, :] / temperature
515
+
516
+ # Top-k filtering
517
+ if top_k > 0:
518
+ indices_to_remove = logits < torch.topk(logits, top_k)[0][..., -1, None]
519
+ logits[indices_to_remove] = -float('Inf')
520
+
521
+ # Top-p (nucleus) filtering
522
+ if top_p < 1.0:
523
+ sorted_logits, sorted_indices = torch.sort(logits, descending=True)
524
+ cumulative_probs = torch.cumsum(F.softmax(sorted_logits, dim=-1), dim=-1)
525
+
526
+ sorted_indices_to_remove = cumulative_probs > top_p
527
+ sorted_indices_to_remove[..., 1:] = sorted_indices_to_remove[..., :-1].clone()
528
+ sorted_indices_to_remove[..., 0] = 0
529
+
530
+ indices_to_remove = sorted_indices_to_remove.scatter(
531
+ 1, sorted_indices, sorted_indices_to_remove
532
+ )
533
+ logits[indices_to_remove] = -float('Inf')
534
+
535
+ # Sample next token
536
+ probs = F.softmax(logits, dim=-1)
537
+ next_token = torch.multinomial(probs, num_samples=1)
538
+
539
+ # Append to sequence
540
+ input_ids = torch.cat([input_ids, next_token], dim=-1)
541
+
542
+ return input_ids
543
+
544
+ def store_experience(self, text: str, metadata: Optional[Dict] = None):
545
+ """
546
+ Store a text experience in long-term memory.
547
+
548
+ Args:
549
+ text: Text to store
550
+ metadata: Optional metadata
551
+ """
552
+ # Tokenize text (simplified - in practice, use proper tokenizer)
553
+ # For demonstration, we'll create dummy embeddings
554
+ # In production, use: tokens = tokenizer.encode(text)
555
+
556
+ # Create a representation (in practice, pass through model)
557
+ dummy_input = torch.zeros(1, 10, self.config.hidden_size)
558
+
559
+ with torch.no_grad():
560
+ # Get representation from model
561
+ representation = self.forward(dummy_input, use_long_term_memory=False)
562
+ # Store in long-term memory
563
+ self.layers[0].long_term_memory._store_new_memories(representation['hidden_states'][-1])
564
+
565
+ def query_memory(self, query: str, top_k: int = 5) -> Tuple[torch.Tensor, torch.Tensor]:
566
+ """
567
+ Query long-term memory for relevant information.
568
+
569
+ Args:
570
+ query: Query text
571
+ top_k: Number of results to retrieve
572
+
573
+ Returns:
574
+ retrieved_memories: Retrieved memory representations
575
+ indices: Indices of retrieved memories
576
+ """
577
+ # Create query representation
578
+ dummy_input = torch.zeros(1, 10, self.config.hidden_size)
579
+
580
+ with torch.no_grad():
581
+ query_representation = self.forward(dummy_input, use_long_term_memory=False)
582
+ query_repr = query_representation['hidden_states'][-1].mean(dim=1)
583
+
584
+ # Retrieve from long-term memory
585
+ retrieved, indices = self.layers[0].long_term_memory.retrieve(query_repr, top_k=top_k)
586
+
587
+ return retrieved, indices
588
+
589
+
590
+ class LMCODEConfig:
591
+ """Configuration for LMCODE model."""
592
+
593
+ def __init__(self, vocab_size: int = 50257, hidden_size: int = 512,
594
+ num_layers: int = 6, num_heads: int = 8,
595
+ short_term_memory_size: int = 512,
596
+ long_term_memory_slots: int = 10000):
597
+ self.vocab_size = vocab_size
598
+ self.hidden_size = hidden_size
599
+ self.num_layers = num_layers
600
+ self.num_heads = num_heads
601
+ self.short_term_memory_size = short_term_memory_size
602
+ self.long_term_memory_slots = long_term_memory_slots
603
+
604
+ def to_dict(self) -> Dict:
605
+ return {
606
+ 'vocab_size': self.vocab_size,
607
+ 'hidden_size': self.hidden_size,
608
+ 'num_layers': self.num_layers,
609
+ 'num_heads': self.num_heads,
610
+ 'short_term_memory_size': self.short_term_memory_size,
611
+ 'long_term_memory_slots': self.long_term_memory_slots
612
+ }
613
+
614
+ @classmethod
615
+ def from_dict(cls, config_dict: Dict) -> 'LMCODEConfig':
616
+ return cls(**config_dict)