""" Complete PALADIM Usage Guide Demonstrates the full continual learning capabilities """ from paladim import PALADIM, PALADIMConfig from datasets import load_dataset from torch.utils.data import DataLoader import torch def demo_continual_learning(): """ Demonstrate PALADIM's continual learning capabilities: - Rapid learning with Plastic Memory (LoRA) - Consolidation to prevent catastrophic forgetting - Learning multiple tasks sequentially """ print("="*60) print("PALADIM Continual Learning Demo") print("="*60) # Initialize PALADIM print("\n1. Initializing PALADIM...") config = PALADIMConfig( model_name="distilbert-base-uncased", max_seq_length=128, learning_rate=5e-4, ) model = PALADIM(config) print("✅ PALADIM initialized with Plastic Memory (LoRA)") # Task 1: Sentiment Analysis print("\n2. Learning Task 1: Sentiment Analysis...") print(" Loading IMDB dataset...") dataset1 = load_dataset("imdb", split="train[:500]") # Create simple dataloader def tokenize(examples): return model.tokenizer( examples["text"], padding="max_length", truncation=True, max_length=128 ) tokenized = dataset1.map(tokenize, batched=True) tokenized = tokenized.rename_column("label", "labels") tokenized.set_format("torch", columns=["input_ids", "attention_mask", "labels"]) train_loader1 = DataLoader(tokenized, batch_size=8, shuffle=True) print(" Training with Plastic Memory (fast adaptation)...") model.learn_task(train_loader1, epochs=1, task_id="sentiment") print(" ✅ Task 1 learned!") # Consolidation (prevent forgetting) print("\n3. Consolidating knowledge...") print(" - Computing Fisher Information Matrix (importance weights)") print(" - Protecting important parameters from future changes") model.consolidate_knowledge(train_loader1) print(" ✅ Knowledge consolidated to Stable Core!") # Task 2: Another task (demo - using same data for simplicity) print("\n4. Learning Task 2: Next Task...") print(" PALADIM will now learn without forgetting Task 1") dataset2 = load_dataset("imdb", split="train[500:1000]") tokenized2 = dataset2.map(tokenize, batched=True) tokenized2 = tokenized2.rename_column("label", "labels") tokenized2.set_format("torch", columns=["input_ids", "attention_mask", "labels"]) train_loader2 = DataLoader(tokenized2, batch_size=8, shuffle=True) model.learn_task(train_loader2, epochs=1, task_id="task2") print(" ✅ Task 2 learned without catastrophic forgetting!") print("\n" + "="*60) print("PALADIM Continual Learning Complete!") print("="*60) print("\nKey Features Demonstrated:") print("✅ Plastic Memory (LoRA) - Fast adaptation") print("✅ Consolidation (EWC) - Prevent forgetting") print("✅ Sequential learning - Multiple tasks") print("\nModel saved and ready for inference!") def demo_basic_usage(): """Simple usage without continual learning""" print("\n" + "="*60) print("Basic PALADIM Usage (Single Task)") print("="*60) # Initialize config = PALADIMConfig(model_name="distilbert-base-uncased") model = PALADIM(config) # Load data print("\nLoading dataset...") dataset = load_dataset("imdb", split="train[:100]") def tokenize(examples): return model.tokenizer( examples["text"], padding="max_length", truncation=True, max_length=128 ) tokenized = dataset.map(tokenize, batched=True) tokenized = tokenized.rename_column("label", "labels") tokenized.set_format("torch", columns=["input_ids", "attention_mask", "labels"]) train_loader = DataLoader(tokenized, batch_size=8, shuffle=True) # Train print("Training...") model.learn_task(train_loader, epochs=1, task_id="sentiment") # Inference print("\nTesting inference...") test_texts = [ "This movie is fantastic!", "Terrible film, very disappointed.", ] for text in test_texts: inputs = model.tokenizer(text, return_tensors="pt", padding=True, truncation=True) with torch.no_grad(): outputs = model.model(**inputs) pred = torch.argmax(outputs.logits, dim=-1).item() sentiment = "Positive" if pred == 1 else "Negative" print(f"'{text}' → {sentiment}") print("\n✅ Basic usage complete!") if __name__ == "__main__": print("Choose demo mode:") print("1. Full Continual Learning (recommended)") print("2. Basic Single-Task Usage") # Run continual learning demo try: demo_continual_learning() except Exception as e: print(f"\nError in continual learning demo: {e}") print("Falling back to basic usage...\n") demo_basic_usage()