alternative_methods.py 23 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617
  1. """
  2. Alternative training methods for voice recognition and wakeword detection.
  3. This module implements few-shot learning, transfer learning, and other advanced training
  4. methodologies inspired by Porcupine and other efficient approaches that require minimal data.
  5. """
  6. import os
  7. import json
  8. import math
  9. from pathlib import Path
  10. from typing import Dict, List, Any, Optional, Tuple, Union
  11. from enum import Enum
  12. import numpy as np
  13. import torch
  14. import torch.nn as nn
  15. import torch.nn.functional as F
  16. from torch.utils.data import DataLoader, Dataset, Subset
  17. import torch.optim as optim
  18. from .base import BaseTrainer, TrainerConfig
  19. from .data_pipeline import AudioProcessingConfig
  20. class TrainingMethod(Enum):
  21. """Enumeration of available training methods."""
  22. STANDARD = "standard" # Traditional full-dataset training
  23. FEW_SHOT = "few_shot" # Few-shot learning with 3-10 samples
  24. META_LEARNING = "meta_learning" # Model-agnostic meta-learning
  25. TRANSFER_LEARNING = "transfer_learning" # Pre-trained backbone + fine-tuning
  26. SELF_SUPERVISED = "self_supervised" # Self-supervised pre-training
  27. DATA_AUGMENTATION = "data_augmentation" # Heavy augmentation with minimal data
  28. CONTRASTIVE = "contrastive" # Contrastive learning approach
  29. PROTOTYPICAL = "prototypical" # Prototypical networks
  30. SIAMESE = "siamese" # Siamese network training
  31. class FewShotConfig:
  32. """Configuration for few-shot learning approaches."""
  33. def __init__(self,
  34. n_way: int = 5, # Number of classes per episode
  35. k_shot: int = 3, # Number of samples per class
  36. n_query: int = 5, # Number of query samples per class
  37. n_episodes: int = 1000, # Number of training episodes
  38. inner_lr: float = 0.01, # Inner loop learning rate
  39. outer_lr: float = 0.001, # Outer loop learning rate
  40. adaptation_steps: int = 5): # Number of adaptation steps
  41. self.n_way = n_way
  42. self.k_shot = k_shot
  43. self.n_query = n_query
  44. self.n_episodes = n_episodes
  45. self.inner_lr = inner_lr
  46. self.outer_lr = outer_lr
  47. self.adaptation_steps = adaptation_steps
  48. class DataAugmentationTrainer:
  49. """
  50. Heavy data augmentation approach similar to Porcupine.
  51. Uses minimal samples (3-5 per class) but applies extensive augmentation:
  52. - Background noise injection
  53. - Speed/pitch variations
  54. - Volume changes
  55. - Temporal shifts
  56. - Spectral augmentation
  57. - Mixup/CutMix
  58. """
  59. def __init__(self, base_trainer: BaseTrainer, augmentation_factor: int = 100):
  60. """
  61. Initialize data augmentation trainer.
  62. Args:
  63. base_trainer: Base trainer instance
  64. augmentation_factor: How many augmented samples per original sample
  65. """
  66. self.base_trainer = base_trainer
  67. self.augmentation_factor = augmentation_factor
  68. self.logger = base_trainer.logger
  69. def create_augmented_dataset(self, original_data: List[Tuple], target_dir: str) -> str:
  70. """
  71. Create heavily augmented dataset from minimal samples.
  72. Args:
  73. original_data: List of (audio_path, label) tuples
  74. target_dir: Directory to save augmented data
  75. Returns:
  76. Path to augmented dataset
  77. """
  78. target_path = Path(target_dir)
  79. target_path.mkdir(parents=True, exist_ok=True)
  80. augmented_samples = []
  81. for audio_path, label in original_data:
  82. # Load original audio
  83. audio = self._load_audio(audio_path)
  84. # Generate augmented versions
  85. for i in range(self.augmentation_factor):
  86. augmented_audio = self._apply_augmentation(audio, intensity=0.8)
  87. # Save augmented sample
  88. aug_filename = f"{Path(audio_path).stem}_aug_{i:03d}.wav"
  89. aug_path = target_path / label / aug_filename
  90. aug_path.parent.mkdir(parents=True, exist_ok=True)
  91. self._save_audio(augmented_audio, aug_path)
  92. augmented_samples.append((str(aug_path), label))
  93. # Save augmentation metadata
  94. metadata = {
  95. 'original_samples': len(original_data),
  96. 'augmented_samples': len(augmented_samples),
  97. 'augmentation_factor': self.augmentation_factor,
  98. 'augmentation_config': self._get_augmentation_config()
  99. }
  100. with open(target_path / "augmentation_metadata.json", 'w') as f:
  101. json.dump(metadata, f, indent=2)
  102. self.logger.info(f"Generated {len(augmented_samples)} augmented samples from {len(original_data)} originals")
  103. return str(target_path)
  104. def _apply_augmentation(self, audio: np.ndarray, intensity: float = 0.5) -> np.ndarray:
  105. """Apply comprehensive audio augmentation."""
  106. augmented = audio.copy()
  107. # Speed variation (0.8x to 1.2x)
  108. if np.random.random() < 0.7:
  109. speed_factor = np.random.uniform(0.8, 1.2)
  110. augmented = self._change_speed(augmented, speed_factor)
  111. # Pitch variation (±2 semitones)
  112. if np.random.random() < 0.6:
  113. pitch_shift = np.random.uniform(-2, 2)
  114. augmented = self._change_pitch(augmented, pitch_shift)
  115. # Volume variation (0.5x to 1.5x)
  116. if np.random.random() < 0.8:
  117. volume_factor = np.random.uniform(0.5, 1.5)
  118. augmented = augmented * volume_factor
  119. # Background noise injection
  120. if np.random.random() < 0.6:
  121. noise_level = np.random.uniform(0.01, 0.1) * intensity
  122. noise = np.random.normal(0, noise_level, augmented.shape)
  123. augmented = augmented + noise
  124. # Temporal shifting
  125. if np.random.random() < 0.5:
  126. shift_samples = int(np.random.uniform(-0.1, 0.1) * len(augmented))
  127. augmented = np.roll(augmented, shift_samples)
  128. # Random cropping and padding
  129. if np.random.random() < 0.4:
  130. crop_ratio = np.random.uniform(0.8, 1.0)
  131. crop_length = int(len(augmented) * crop_ratio)
  132. start_idx = np.random.randint(0, len(augmented) - crop_length)
  133. augmented = augmented[start_idx:start_idx + crop_length]
  134. # Pad back to original length
  135. pad_length = len(audio) - len(augmented)
  136. if pad_length > 0:
  137. pad_left = np.random.randint(0, pad_length + 1)
  138. pad_right = pad_length - pad_left
  139. augmented = np.pad(augmented, (pad_left, pad_right), mode='constant')
  140. # Ensure audio stays in valid range
  141. augmented = np.clip(augmented, -1.0, 1.0)
  142. return augmented
  143. def _change_speed(self, audio: np.ndarray, factor: float) -> np.ndarray:
  144. """Change audio speed using resampling."""
  145. # Simple linear interpolation for speed change
  146. indices = np.arange(0, len(audio), factor)
  147. indices = indices[indices < len(audio)]
  148. return np.interp(indices, np.arange(len(audio)), audio)
  149. def _change_pitch(self, audio: np.ndarray, semitones: float) -> np.ndarray:
  150. """Change pitch by shifting in frequency domain (simplified)."""
  151. # This is a simplified pitch shift - in practice, you'd use librosa or similar
  152. factor = 2 ** (semitones / 12.0)
  153. return self._change_speed(audio, factor)
  154. def _load_audio(self, path: str) -> np.ndarray:
  155. """Load audio file (simplified - implement with librosa/soundfile)."""
  156. # Placeholder - implement actual audio loading
  157. return np.random.randn(16000) # 1 second at 16kHz
  158. def _save_audio(self, audio: np.ndarray, path: str):
  159. """Save audio file (simplified - implement with librosa/soundfile)."""
  160. # Placeholder - implement actual audio saving
  161. pass
  162. def _get_augmentation_config(self) -> Dict[str, Any]:
  163. """Get augmentation configuration for metadata."""
  164. return {
  165. 'speed_variation': {'min': 0.8, 'max': 1.2, 'probability': 0.7},
  166. 'pitch_variation': {'min': -2, 'max': 2, 'probability': 0.6},
  167. 'volume_variation': {'min': 0.5, 'max': 1.5, 'probability': 0.8},
  168. 'noise_injection': {'min': 0.01, 'max': 0.1, 'probability': 0.6},
  169. 'temporal_shift': {'min': -0.1, 'max': 0.1, 'probability': 0.5},
  170. 'random_crop': {'min_ratio': 0.8, 'max_ratio': 1.0, 'probability': 0.4}
  171. }
  172. class PrototypicalNetworkTrainer:
  173. """
  174. Prototypical Networks for few-shot learning.
  175. Based on "Prototypical Networks for Few-shot Learning" by Snell et al.
  176. Creates class prototypes from support samples and classifies query samples
  177. based on distance to prototypes.
  178. """
  179. def __init__(self, model: nn.Module, config: FewShotConfig, device: str):
  180. """
  181. Initialize prototypical network trainer.
  182. Args:
  183. model: Embedding model (backbone)
  184. config: Few-shot learning configuration
  185. device: Training device
  186. """
  187. self.model = model
  188. self.config = config
  189. self.device = device
  190. def train_episode(self, support_data: torch.Tensor, support_labels: torch.Tensor,
  191. query_data: torch.Tensor, query_labels: torch.Tensor) -> float:
  192. """
  193. Train on a single episode using prototypical networks.
  194. Args:
  195. support_data: Support set data (n_way * k_shot, ...)
  196. support_labels: Support set labels
  197. query_data: Query set data (n_way * n_query, ...)
  198. query_labels: Query set labels
  199. Returns:
  200. Episode loss
  201. """
  202. self.model.train()
  203. # Compute embeddings
  204. support_embeddings = self.model(support_data)
  205. query_embeddings = self.model(query_data)
  206. # Compute prototypes (class centroids)
  207. n_way = self.config.n_way
  208. k_shot = self.config.k_shot
  209. prototypes = []
  210. for class_idx in range(n_way):
  211. class_embeddings = support_embeddings[class_idx * k_shot:(class_idx + 1) * k_shot]
  212. prototype = torch.mean(class_embeddings, dim=0)
  213. prototypes.append(prototype)
  214. prototypes = torch.stack(prototypes) # (n_way, embedding_dim)
  215. # Compute distances from query embeddings to prototypes
  216. distances = self._compute_distances(query_embeddings, prototypes)
  217. # Convert distances to logits (negative distances)
  218. logits = -distances
  219. # Compute loss
  220. loss = F.cross_entropy(logits, query_labels)
  221. return loss.item()
  222. def _compute_distances(self, query_embeddings: torch.Tensor,
  223. prototypes: torch.Tensor) -> torch.Tensor:
  224. """
  225. Compute Euclidean distances between query embeddings and prototypes.
  226. Args:
  227. query_embeddings: Query embeddings (n_query, embedding_dim)
  228. prototypes: Class prototypes (n_way, embedding_dim)
  229. Returns:
  230. Distance matrix (n_query, n_way)
  231. """
  232. n_query = query_embeddings.size(0)
  233. n_way = prototypes.size(0)
  234. # Expand dimensions for broadcasting
  235. query_expanded = query_embeddings.unsqueeze(1).expand(n_query, n_way, -1)
  236. prototype_expanded = prototypes.unsqueeze(0).expand(n_query, n_way, -1)
  237. # Compute Euclidean distances
  238. distances = torch.pow(query_expanded - prototype_expanded, 2).sum(dim=2)
  239. return distances
  240. class MAMLTrainer:
  241. """
  242. Model-Agnostic Meta-Learning (MAML) trainer.
  243. Based on "Model-Agnostic Meta-Learning for Fast Adaptation of Deep Networks"
  244. by Finn et al. Learns initialization parameters that can quickly adapt to new tasks.
  245. """
  246. def __init__(self, model: nn.Module, config: FewShotConfig, device: str):
  247. """
  248. Initialize MAML trainer.
  249. Args:
  250. model: Model to meta-learn
  251. config: Few-shot learning configuration
  252. device: Training device
  253. """
  254. self.model = model
  255. self.config = config
  256. self.device = device
  257. self.meta_optimizer = optim.Adam(self.model.parameters(), lr=config.outer_lr)
  258. def meta_train_step(self, tasks: List[Tuple]) -> float:
  259. """
  260. Perform one meta-training step across multiple tasks.
  261. Args:
  262. tasks: List of (support_data, support_labels, query_data, query_labels) tuples
  263. Returns:
  264. Meta-loss across all tasks
  265. """
  266. self.meta_optimizer.zero_grad()
  267. meta_loss = 0.0
  268. for task_data in tasks:
  269. support_data, support_labels, query_data, query_labels = task_data
  270. # Create a copy of the model for this task
  271. fast_weights = self._get_model_params()
  272. # Inner loop: adapt to support set
  273. for _ in range(self.config.adaptation_steps):
  274. support_loss = self._compute_loss(support_data, support_labels, fast_weights)
  275. # Compute gradients with respect to fast weights
  276. grads = torch.autograd.grad(support_loss, fast_weights, create_graph=True)
  277. # Update fast weights
  278. fast_weights = [w - self.config.inner_lr * g for w, g in zip(fast_weights, grads)]
  279. # Outer loop: compute meta-loss on query set
  280. query_loss = self._compute_loss(query_data, query_labels, fast_weights)
  281. meta_loss += query_loss
  282. # Average meta-loss
  283. meta_loss = meta_loss / len(tasks)
  284. # Backpropagation for meta-parameters
  285. meta_loss.backward()
  286. self.meta_optimizer.step()
  287. return meta_loss.item()
  288. def _get_model_params(self) -> List[torch.Tensor]:
  289. """Get model parameters as a list."""
  290. return [p.clone() for p in self.model.parameters()]
  291. def _compute_loss(self, data: torch.Tensor, labels: torch.Tensor,
  292. weights: List[torch.Tensor]) -> torch.Tensor:
  293. """
  294. Compute loss using specific weights.
  295. Args:
  296. data: Input data
  297. labels: Target labels
  298. weights: Model weights to use
  299. Returns:
  300. Loss value
  301. """
  302. # This is a simplified version - actual implementation would need to
  303. # substitute weights into the model forward pass
  304. outputs = self.model(data) # Would use weights parameter
  305. return F.cross_entropy(outputs, labels)
  306. class TransferLearningTrainer:
  307. """
  308. Transfer learning trainer using pre-trained backbones.
  309. Uses pre-trained models (e.g., from speech recognition, speaker verification)
  310. and fine-tunes for specific tasks with minimal data.
  311. """
  312. def __init__(self, backbone_path: str, num_classes: int, device: str):
  313. """
  314. Initialize transfer learning trainer.
  315. Args:
  316. backbone_path: Path to pre-trained backbone model
  317. num_classes: Number of target classes
  318. device: Training device
  319. """
  320. self.backbone_path = backbone_path
  321. self.num_classes = num_classes
  322. self.device = device
  323. self.model = None
  324. def create_model(self, freeze_backbone: bool = True) -> nn.Module:
  325. """
  326. Create transfer learning model.
  327. Args:
  328. freeze_backbone: Whether to freeze backbone parameters
  329. Returns:
  330. Transfer learning model
  331. """
  332. # Load pre-trained backbone
  333. backbone = self._load_backbone(self.backbone_path)
  334. # Freeze backbone if requested
  335. if freeze_backbone:
  336. for param in backbone.parameters():
  337. param.requires_grad = False
  338. # Add classification head
  339. feature_dim = self._get_backbone_feature_dim(backbone)
  340. classifier = nn.Sequential(
  341. nn.Dropout(0.5),
  342. nn.Linear(feature_dim, 256),
  343. nn.ReLU(),
  344. nn.Dropout(0.3),
  345. nn.Linear(256, self.num_classes)
  346. )
  347. # Combine backbone and classifier
  348. self.model = nn.Sequential(backbone, classifier)
  349. return self.model
  350. def fine_tune(self, data_loader: DataLoader, num_epochs: int = 50) -> Dict[str, List[float]]:
  351. """
  352. Fine-tune the model on target data.
  353. Args:
  354. data_loader: Target dataset loader
  355. num_epochs: Number of fine-tuning epochs
  356. Returns:
  357. Training metrics
  358. """
  359. optimizer = optim.Adam(self.model.parameters(), lr=0.0001)
  360. criterion = nn.CrossEntropyLoss()
  361. losses = []
  362. accuracies = []
  363. for epoch in range(num_epochs):
  364. epoch_loss = 0.0
  365. epoch_correct = 0
  366. epoch_total = 0
  367. for data, targets in data_loader:
  368. data, targets = data.to(self.device), targets.to(self.device)
  369. optimizer.zero_grad()
  370. outputs = self.model(data)
  371. loss = criterion(outputs, targets)
  372. loss.backward()
  373. optimizer.step()
  374. epoch_loss += loss.item()
  375. _, predicted = torch.max(outputs.data, 1)
  376. epoch_total += targets.size(0)
  377. epoch_correct += (predicted == targets).sum().item()
  378. epoch_loss /= len(data_loader)
  379. epoch_acc = epoch_correct / epoch_total
  380. losses.append(epoch_loss)
  381. accuracies.append(epoch_acc)
  382. return {'losses': losses, 'accuracies': accuracies}
  383. def _load_backbone(self, path: str) -> nn.Module:
  384. """Load pre-trained backbone model."""
  385. # Placeholder - implement actual model loading
  386. # This would load a pre-trained model and remove the classification head
  387. return nn.Sequential()
  388. def _get_backbone_feature_dim(self, backbone: nn.Module) -> int:
  389. """Get feature dimension of backbone output."""
  390. # Placeholder - determine feature dimension
  391. return 512
  392. class AlternativeTrainingManager:
  393. """
  394. Manager for alternative training methods.
  395. Provides a unified interface for selecting and configuring different
  396. training approaches based on available data and requirements.
  397. """
  398. def __init__(self, base_config: TrainerConfig):
  399. """
  400. Initialize alternative training manager.
  401. Args:
  402. base_config: Base trainer configuration
  403. """
  404. self.base_config = base_config
  405. self.training_method = TrainingMethod.STANDARD
  406. self.few_shot_config = FewShotConfig()
  407. def set_training_method(self, method: TrainingMethod, **kwargs):
  408. """
  409. Set the training method and its configuration.
  410. Args:
  411. method: Training method to use
  412. **kwargs: Method-specific configuration
  413. """
  414. self.training_method = method
  415. if method == TrainingMethod.FEW_SHOT:
  416. self.few_shot_config = FewShotConfig(**kwargs)
  417. elif method == TrainingMethod.DATA_AUGMENTATION:
  418. self.augmentation_factor = kwargs.get('augmentation_factor', 100)
  419. # Add other method configurations as needed
  420. def get_recommended_method(self, dataset_size: int, samples_per_class: int) -> TrainingMethod:
  421. """
  422. Recommend training method based on available data.
  423. Args:
  424. dataset_size: Total dataset size
  425. samples_per_class: Average samples per class
  426. Returns:
  427. Recommended training method
  428. """
  429. if samples_per_class <= 5:
  430. if dataset_size < 100:
  431. return TrainingMethod.DATA_AUGMENTATION
  432. else:
  433. return TrainingMethod.FEW_SHOT
  434. elif samples_per_class <= 20:
  435. return TrainingMethod.TRANSFER_LEARNING
  436. elif samples_per_class <= 50:
  437. return TrainingMethod.SELF_SUPERVISED
  438. else:
  439. return TrainingMethod.STANDARD
  440. def create_trainer(self, model: nn.Module, device: str) -> Any:
  441. """
  442. Create appropriate trainer based on selected method.
  443. Args:
  444. model: Model to train
  445. device: Training device
  446. Returns:
  447. Trainer instance for the selected method
  448. """
  449. if self.training_method == TrainingMethod.FEW_SHOT:
  450. return PrototypicalNetworkTrainer(model, self.few_shot_config, device)
  451. elif self.training_method == TrainingMethod.META_LEARNING:
  452. return MAMLTrainer(model, self.few_shot_config, device)
  453. elif self.training_method == TrainingMethod.TRANSFER_LEARNING:
  454. backbone_path = self.base_config.custom_params.get('backbone_path', '')
  455. num_classes = self.base_config.custom_params.get('num_classes', 10)
  456. return TransferLearningTrainer(backbone_path, num_classes, device)
  457. elif self.training_method == TrainingMethod.DATA_AUGMENTATION:
  458. base_trainer = BaseTrainer(self.base_config) # This would be the actual trainer
  459. return DataAugmentationTrainer(base_trainer, self.augmentation_factor)
  460. else:
  461. # Return standard trainer
  462. return BaseTrainer(self.base_config)
  463. def get_training_config(self) -> Dict[str, Any]:
  464. """Get configuration for the selected training method."""
  465. config = {
  466. 'method': self.training_method.value,
  467. 'base_config': self.base_config.to_dict()
  468. }
  469. if self.training_method in [TrainingMethod.FEW_SHOT, TrainingMethod.META_LEARNING]:
  470. config['few_shot_config'] = {
  471. 'n_way': self.few_shot_config.n_way,
  472. 'k_shot': self.few_shot_config.k_shot,
  473. 'n_query': self.few_shot_config.n_query,
  474. 'n_episodes': self.few_shot_config.n_episodes,
  475. 'inner_lr': self.few_shot_config.inner_lr,
  476. 'outer_lr': self.few_shot_config.outer_lr,
  477. 'adaptation_steps': self.few_shot_config.adaptation_steps
  478. }
  479. return config
  480. def create_alternative_training_manager(base_config: TrainerConfig) -> AlternativeTrainingManager:
  481. """
  482. Factory function to create alternative training manager.
  483. Args:
  484. base_config: Base trainer configuration
  485. Returns:
  486. AlternativeTrainingManager instance
  487. """
  488. return AlternativeTrainingManager(base_config)