lm_memory_code / run_demo.py
userkuku's picture
Upload run_demo.py
62a3dea verified
Raw
History Blame Contribute Delete
10.9 kB
"""
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)
# Create a small model for demonstration
config = LMCODEConfig(
vocab_size=1000, # Smaller vocab for demo
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}")
# Count parameters
total_params = sum(p.numel() for p in model.parameters())
print(f"\nTotal parameters: {total_params:,}")
# Forward pass
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']}")
# Test generation
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)
# Store some experiences
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]}...")
# Try to retrieve
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()}")
# Consolidate memories
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)
# Create model
config = LMCODEConfig(
vocab_size=1000,
hidden_size=64, # Small for fast demo
num_layers=2,
num_heads=4,
short_term_memory_size=64,
long_term_memory_slots=500
)
model = LMCODE(config)
# Create dataset
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)}")
# Create trainer
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)
# Train for 3 epochs
print("\nTraining model (3 epochs, small for demo)...")
history = trainer.train(
train_dataset,
num_epochs=3,
batch_size=16,
eval_dataset=eval_dataset
)
# Show results
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}")
# Save checkpoint
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)
# Create test sequences
test_sequences = []
for _ in range(10):
seq = torch.randint(0, model.config.vocab_size, (1, 15))
test_sequences.append(seq)
# Analyze memory capacity
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%}")
# Compute efficiency
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}")
# Generate report
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)
# Create sample input
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})")
# Create training history plot
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):
# Create dummy batch
input_ids = torch.randint(0, 1000, (4, 20))
labels = torch.randint(0, 1000, (4, 20))
# Forward pass
with torch.no_grad():
outputs = model(input_ids, labels=labels, use_long_term_memory=True)
# Record step
monitor.record_step(step, outputs)
# Get statistics
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)
# Demo 1: Basic usage
model, config = demo_basic_usage()
# Demo 2: Memory operations
demo_memory_operations(model)
# Demo 3: Training
model, history = demo_training()
# Demo 4: Memory analysis
demo_memory_analysis(model)
# Demo 5: Visualization
demo_visualization(model)
# Demo 6: Monitoring
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()