|
| 1 | +//! MLP Router example - showing how to use neural network for routing |
| 2 | +//! |
| 3 | +//! Run with: cargo run --example mlp_router |
| 4 | +
|
| 5 | +use mobile_ai_orchestrator::mlp::MLP; |
| 6 | +use mobile_ai_orchestrator::reservoir::encode_text; |
| 7 | + |
| 8 | +fn main() { |
| 9 | + println!("MLP Router Example\n"); |
| 10 | + |
| 11 | + // Create MLP for routing decisions |
| 12 | + // Input: 384-dim text encoding |
| 13 | + // Hidden: 100 → 50 neurons |
| 14 | + // Output: 3 classes (Local, Remote, Hybrid) |
| 15 | + let mlp = MLP::new(384, vec![100, 50], 3); |
| 16 | + |
| 17 | + println!("=== MLP Architecture ==="); |
| 18 | + println!("Input size: {}", mlp.input_size()); |
| 19 | + println!("Output size: {}", mlp.output_size()); |
| 20 | + println!("Hidden layers: 100 → 50"); |
| 21 | + |
| 22 | + // Example queries |
| 23 | + let queries = vec![ |
| 24 | + "How do I iterate a HashMap?", |
| 25 | + "Can you formally prove this theorem?", |
| 26 | + "Help me debug this complex multi-threaded race condition", |
| 27 | + ]; |
| 28 | + |
| 29 | + println!("\n=== Routing Decisions ==="); |
| 30 | + for query in queries { |
| 31 | + // Encode query |
| 32 | + let encoding = encode_text(query, 384); |
| 33 | + |
| 34 | + // Forward through MLP |
| 35 | + let scores = mlp.forward(&encoding); |
| 36 | + |
| 37 | + // Apply softmax to get probabilities |
| 38 | + let probs = MLP::softmax(&scores); |
| 39 | + |
| 40 | + // Get decision |
| 41 | + let decision = MLP::argmax(&probs); |
| 42 | + |
| 43 | + let labels = ["Local", "Remote", "Hybrid"]; |
| 44 | + println!("\nQuery: '{}'", query); |
| 45 | + println!("Probabilities:"); |
| 46 | + for (i, (label, prob)) in labels.iter().zip(&probs).enumerate() { |
| 47 | + let marker = if i == decision { "→" } else { " " }; |
| 48 | + println!(" {} {}: {:.3}", marker, label, prob); |
| 49 | + } |
| 50 | + println!("Decision: {}", labels[decision]); |
| 51 | + } |
| 52 | + |
| 53 | + // Example: Training step (simplified) |
| 54 | + println!("\n=== Training Example ==="); |
| 55 | + let mut trainable_mlp = MLP::new(10, vec![20], 3); |
| 56 | + let input = vec![1.0; 10]; |
| 57 | + let target = vec![0.0, 1.0, 0.0]; // Correct answer: Remote |
| 58 | + |
| 59 | + let loss_before = trainable_mlp.train_step(&input, &target, 0.01); |
| 60 | + println!("Loss before training: {:.4}", loss_before); |
| 61 | + |
| 62 | + // Train for a few steps |
| 63 | + for _ in 0..100 { |
| 64 | + trainable_mlp.train_step(&input, &target, 0.01); |
| 65 | + } |
| 66 | + |
| 67 | + let loss_after = trainable_mlp.train_step(&input, &target, 0.01); |
| 68 | + println!("Loss after 100 steps: {:.4}", loss_after); |
| 69 | + println!("Improvement: {:.4}", loss_before - loss_after); |
| 70 | + |
| 71 | + println!("\n✅ MLP router example completed!"); |
| 72 | + println!("\nNote: In production, train on real user feedback data"); |
| 73 | + println!("Collect: (query, user-corrected routing decision)"); |
| 74 | + println!("Train: offline, deploy weights via model update"); |
| 75 | +} |
0 commit comments