| """ |
| Demo script for LMCODE (Language Model with Memory CODE). |
| |
| Demonstrates the dual memory system in action. |
| """ |
|
|
| import torch |
| import matplotlib.pyplot as plt |
| from model_architecture import LMCODE, LMCODEConfig |
| from training import MemoryAwareTrainer, MemoryDataset, create_synthetic_dataset |
| from utils import ( |
| analyze_memory_capacity, |
| compute_memory_efficiency, |
| visualize_memory_flow, |
| plot_training_history, |
| MemoryMonitor, |
| generate_memory_report |
| ) |
| import numpy as np |
|
|
|
|
| def demo_basic_usage(): |
| """Demonstrate basic model usage.""" |
| print("=" * 60) |
| print("DEMO 1: Basic Model Usage") |
| print("=" * 60) |
| |
| |
| config = LMCODEConfig( |
| vocab_size=1000, |
| hidden_size=128, |
| num_layers=3, |
| num_heads=4, |
| short_term_memory_size=128, |
| long_term_memory_slots=1000 |
| ) |
| |
| model = LMCODE(config) |
| |
| print(f"\nModel created with config:") |
| print(f" Vocabulary size: {config.vocab_size}") |
| print(f" Hidden size: {config.hidden_size}") |
| print(f" Number of layers: {config.num_layers}") |
| print(f" Number of heads: {config.num_heads}") |
| print(f" Short-term memory size: {config.short_term_memory_size}") |
| print(f" Long-term memory slots: {config.long_term_memory_slots}") |
| |
| |
| total_params = sum(p.numel() for p in model.parameters()) |
| print(f"\nTotal parameters: {total_params:,}") |
| |
| |
| batch_size = 2 |
| seq_len = 20 |
| input_ids = torch.randint(0, config.vocab_size, (batch_size, seq_len)) |
| |
| print(f"\nInput shape: {input_ids.shape}") |
| |
| with torch.no_grad(): |
| outputs = model(input_ids, use_long_term_memory=True) |
| |
| print(f"Output logits shape: {outputs['logits'].shape}") |
| print(f"Loss: {outputs['loss']}") |
| |
| |
| print("\nTesting text generation...") |
| start_tokens = torch.randint(0, config.vocab_size, (1, 5)) |
| |
| with torch.no_grad(): |
| generated = model.generate( |
| start_tokens, |
| max_length=30, |
| temperature=1.0, |
| top_k=50, |
| top_p=0.9, |
| use_long_term_memory=True |
| ) |
| |
| print(f"Input tokens: {start_tokens[0].tolist()}") |
| print(f"Generated tokens (first 20): {generated[0][:20].tolist()}") |
| print(f"Generated shape: {generated.shape}") |
| |
| return model, config |
|
|
|
|
| def demo_memory_operations(model): |
| """Demonstrate memory store and retrieve operations.""" |
| print("\n" + "=" * 60) |
| print("DEMO 2: Memory Store and Retrieve Operations") |
| print("=" * 60) |
| |
| |
| experiences = [ |
| "The quick brown fox jumps over the lazy dog", |
| "Machine learning is a subset of artificial intelligence", |
| "Python is a popular programming language for data science", |
| "Neural networks can learn complex patterns", |
| "Transformers have revolutionized natural language processing" |
| ] |
| |
| print("\nStoring experiences in long-term memory...") |
| for exp in experiences: |
| model.store_experience(exp) |
| print(f" Stored: {exp[:50]}...") |
| |
| |
| print("\nQuerying memory...") |
| queries = [ |
| "programming language", |
| "neural networks", |
| "machine learning" |
| ] |
| |
| for query in queries: |
| retrieved, indices = model.query_memory(query, top_k=3) |
| print(f"\nQuery: '{query}'") |
| print(f" Retrieved shape: {retrieved.shape}") |
| print(f" Top indices: {indices[0].tolist()}") |
| |
| |
| print("\nConsolidating memories (merging similar ones)...") |
| for i, layer in enumerate(model.layers): |
| before_active = (torch.sigmoid(layer.long_term_memory.memory_importance) > 0.1).sum().item() |
| layer.long_term_memory.consolidate_memories(threshold=0.9) |
| after_active = (torch.sigmoid(layer.long_term_memory.memory_importance) > 0.1).sum().item() |
| print(f" Layer {i}: {before_active} -> {after_active} active memories") |
|
|
|
|
| def demo_training(): |
| """Demonstrate training with memory-aware trainer.""" |
| print("\n" + "=" * 60) |
| print("DEMO 3: Training with Memory-Aware Trainer") |
| print("=" * 60) |
| |
| |
| config = LMCODEConfig( |
| vocab_size=1000, |
| hidden_size=64, |
| num_layers=2, |
| num_heads=4, |
| short_term_memory_size=64, |
| long_term_memory_slots=500 |
| ) |
| |
| model = LMCODE(config) |
| |
| |
| print("\nCreating synthetic dataset...") |
| train_data = create_synthetic_dataset(num_samples=200, seq_len=20, vocab_size=1000) |
| train_dataset = MemoryDataset(train_data, memory_sample_ratio=0.2) |
| |
| eval_data = create_synthetic_dataset(num_samples=50, seq_len=20, vocab_size=1000) |
| eval_dataset = MemoryDataset(eval_data, memory_sample_ratio=0.2) |
| |
| print(f"Training samples: {len(train_data)}") |
| print(f"Evaluation samples: {len(eval_data)}") |
| |
| |
| trainer_config = { |
| 'learning_rate': 1e-3, |
| 'weight_decay': 0.01, |
| 'gradient_clip': 1.0, |
| 'memory_consolidation_interval': 20, |
| 'warmup_steps': 5, |
| 'total_steps': 200 |
| } |
| |
| trainer = MemoryAwareTrainer(model, trainer_config) |
| |
| |
| print("\nTraining model (3 epochs, small for demo)...") |
| history = trainer.train( |
| train_dataset, |
| num_epochs=3, |
| batch_size=16, |
| eval_dataset=eval_dataset |
| ) |
| |
| |
| print("\nTraining complete!") |
| print(f"Final train loss: {history['train_loss'][-1]:.4f}") |
| if history['eval_loss']: |
| print(f"Final eval loss: {history['eval_loss'][-1]:.4f}") |
| |
| |
| trainer.save_checkpoint('demo_model.pt') |
| |
| return model, history |
|
|
|
|
| def demo_memory_analysis(model): |
| """Demonstrate memory analysis tools.""" |
| print("\n" + "=" * 60) |
| print("DEMO 4: Memory Analysis and Efficiency") |
| print("=" * 60) |
| |
| |
| test_sequences = [] |
| for _ in range(10): |
| seq = torch.randint(0, model.config.vocab_size, (1, 15)) |
| test_sequences.append(seq) |
| |
| |
| print("\nAnalyzing memory capacity...") |
| analysis = analyze_memory_capacity(model, test_sequences) |
| |
| print(f"Total memories stored: {analysis['total_memories']}") |
| print(f"Successful retrievals: {analysis['successful_retrievals']}") |
| print(f"Average similarity: {analysis.get('average_similarity', 'N/A')}") |
| print(f"Capacity utilization: {analysis['capacity_utilization']:.2%}") |
| |
| |
| print("\nComputing memory efficiency...") |
| efficiency = compute_memory_efficiency(model) |
| |
| print(f"Total parameters: {efficiency['total_parameters']:,}") |
| print(f"Memory parameters: {efficiency['memory_parameters']:,}") |
| print(f"Memory parameter ratio: {efficiency['memory_parameter_ratio']:.2%}") |
| print(f"Total memory slots: {efficiency['total_memory_slots']}") |
| print(f"Parameters per slot: {efficiency['parameters_per_memory_slot']:.1f}") |
| |
| |
| print("\nGenerating memory report...") |
| report = generate_memory_report(model, test_sequences, 'demo_memory_report.json') |
| print(f"Report saved to demo_memory_report.json") |
| |
| return analysis, efficiency |
|
|
|
|
| def demo_visualization(model): |
| """Demonstrate visualization tools.""" |
| print("\n" + "=" * 60) |
| print("DEMO 5: Visualization Tools") |
| print("=" * 60) |
| |
| |
| input_seq = torch.randint(0, model.config.vocab_size, (1, 25)) |
| |
| print("\nGenerating memory flow visualization...") |
| try: |
| fig = visualize_memory_flow(model, input_seq.squeeze(0)) |
| plt.savefig('demo_memory_flow.png', dpi=150, bbox_inches='tight') |
| plt.close() |
| print("Saved to demo_memory_flow.png") |
| except Exception as e: |
| print(f"Note: Visualization requires display (error: {e})") |
| |
| |
| print("\nGenerating training history plot...") |
| history = { |
| 'train_loss': [2.5, 2.0, 1.5, 1.2, 1.0, 0.9, 0.85], |
| 'eval_loss': [2.4, 1.9, 1.4, 1.3, 1.1, 1.0, 0.95], |
| 'memory_stats': [ |
| {'layer_0_lt_active_count': i * 10} for i in range(7) |
| ] |
| } |
| |
| try: |
| fig = plot_training_history(history) |
| plt.savefig('demo_training_history.png', dpi=150, bbox_inches='tight') |
| plt.close() |
| print("Saved to demo_training_history.png") |
| except Exception as e: |
| print(f"Note: Visualization requires display (error: {e})") |
|
|
|
|
| def demo_monitor(): |
| """Demonstrate memory monitoring during training.""" |
| print("\n" + "=" * 60) |
| print("DEMO 6: Memory Monitoring") |
| print("=" * 60) |
| |
| config = LMCODEConfig( |
| vocab_size=1000, |
| hidden_size=64, |
| num_layers=2, |
| num_heads=4, |
| short_term_memory_size=64, |
| long_term_memory_slots=500 |
| ) |
| |
| model = LMCODE(config) |
| monitor = MemoryMonitor(model) |
| |
| print("\nSimulating training steps with monitoring...") |
| for step in range(10): |
| |
| input_ids = torch.randint(0, 1000, (4, 20)) |
| labels = torch.randint(0, 1000, (4, 20)) |
| |
| |
| with torch.no_grad(): |
| outputs = model(input_ids, labels=labels, use_long_term_memory=True) |
| |
| |
| monitor.record_step(step, outputs) |
| |
| |
| print("\nMemory monitoring statistics:") |
| stats = monitor.get_statistics() |
| for key, val in stats.items(): |
| print(f" {key}:") |
| print(f" Mean: {val['mean']:.4f}") |
| print(f" Std: {val['std']:.4f}") |
| print(f" Latest: {val['latest']:.4f}") |
| |
| print("\nNote: Full visualization requires display environment") |
|
|
|
|
| def main(): |
| """Run all demos.""" |
| print("\n" + "=" * 60) |
| print("LMCODE: Language Model with Memory CODE - Demo") |
| print("=" * 60) |
| |
| |
| model, config = demo_basic_usage() |
| |
| |
| demo_memory_operations(model) |
| |
| |
| model, history = demo_training() |
| |
| |
| demo_memory_analysis(model) |
| |
| |
| demo_visualization(model) |
| |
| |
| demo_monitor() |
| |
| print("\n" + "=" * 60) |
| print("All demos completed successfully!") |
| print("=" * 60) |
| print("\nGenerated files:") |
| print(" - demo_model.pt (trained model checkpoint)") |
| print(" - demo_memory_report.json (memory analysis)") |
| print(" - demo_memory_flow.png (memory flow visualization)") |
| print(" - demo_training_history.png (training history)") |
|
|
|
|
| if __name__ == '__main__': |
| main() |
|
|