import torch import torch.nn as nn from torch.nn.utils.fusion import fuse_conv_bn_eval from typing import Optional, Dict, Any import numpy as np import time from pathlib import Path from ..config import settings from .rawnet import RawNet, RawNetConfig class QuantizedModelWrapper: """ Wrapper for INT8 quantized RawNet model. Handles dynamic quantization, inference, and model loading/saving. """ def __init__( self, model: Optional[RawNet] = None, config: Optional[RawNetConfig] = None, quantized: bool = settings.quantized, ): self.config = config or RawNetConfig() self.quantized = quantized self._model = model self._quantized_model: Optional[nn.Module] = None self._device = torch.device("cpu") # Force CPU for on-premise self._warmup_done = False @property def model(self) -> nn.Module: if self.quantized and self._quantized_model is not None: return self._quantized_model return self._model def load_model(self, model_path: str = None, strict: bool = True) -> None: """Load model weights from file.""" path = model_path or settings.model_path path = Path(path) if not path.exists(): path.parent.mkdir(parents=True, exist_ok=True) if self._model is None: self._model = RawNet(self.config) self._model.to(self._device) self._model.eval() torch.save(self._model.state_dict(), path) else: if self._model is None: self._model = RawNet(self.config) state_dict = torch.load(path, map_location=self._device, weights_only=True) self._model.load_state_dict(state_dict, strict=strict) self._model.to(self._device) self._model.eval() if self.quantized: self.quantize() def quantize(self) -> None: """Apply Conv-BN fusion and INT8 Linear-only quantization for edge latency.""" if self._model is None: raise RuntimeError("Model not loaded. Call load_model() first.") print("Applying Conv-BN fusion + INT8 Linear quantization...") start = time.time() # 1. Pre-cache sinc filters to eliminate torch.sinc() at runtime self._model.sinc_conv.precompute_kernel() # 2. Fuse BatchNorm into Conv1D layers (eliminates 4 BN forward passes) self._fuse_conv_bn() # 3. Quantize Linear layers only (GRU dequantization overhead # exceeds INT8 compute savings on 2 vCPUs) self._quantized_model = torch.quantization.quantize_dynamic( self._model, qconfig_spec={nn.Linear}, dtype=torch.qint8, inplace=False, ) self._quantized_model.to(self._device) self._quantized_model.eval() elapsed = (time.time() - start) * 1000 print(f"Optimization completed in {elapsed:.1f}ms") # Verify quantization self._verify_quantization() def _fuse_conv_bn(self) -> None: """Fold BatchNorm parameters into Conv1D biases for inference speedup.""" model = self._model for block in getattr(model, 'res_blocks', []): if hasattr(block, 'conv1') and hasattr(block, 'bn1'): block.conv1 = fuse_conv_bn_eval(block.conv1, block.bn1) block.bn1 = nn.Identity() if hasattr(block, 'conv2') and hasattr(block, 'bn2'): block.conv2 = fuse_conv_bn_eval(block.conv2, block.bn2) block.bn2 = nn.Identity() # Detach fused params so they are leaf tensors (required by deepcopy) for conv in [block.conv1, block.conv2]: if hasattr(conv, 'weight'): conv.weight = nn.Parameter(conv.weight.detach().clone()) if hasattr(conv, 'bias') and conv.bias is not None: conv.bias = nn.Parameter(conv.bias.detach().clone()) def _verify_quantization(self) -> None: """Verify that quantization was applied to target layers.""" quantized_layers = 0 total_layers = 0 for name, module in self._quantized_model.named_modules(): if isinstance(module, (nn.Linear, nn.GRU)): total_layers += 1 if hasattr(module, 'weight') and module.weight.dtype == torch.qint8: quantized_layers += 1 print(f"Quantized {quantized_layers}/{total_layers} target layers") def warmup(self, num_runs: int = 3) -> float: """Run warmup inferences to stabilize timing.""" if self._model is None: raise RuntimeError("Model not loaded") model = self.model dummy_input = torch.randn(1, 1, 24000, device=self._device) times = [] with torch.no_grad(): for _ in range(num_runs): start = time.perf_counter() _ = model(dummy_input) times.append((time.perf_counter() - start) * 1000) self._warmup_done = True avg_time = sum(times) / len(times) print(f"Warmup complete. Avg inference: {avg_time:.2f}ms") return avg_time def predict(self, audio: np.ndarray) -> Dict[str, Any]: """ Run inference on preprocessed audio chunk. audio: [T] or [1, T] float32 numpy array Returns dict with logits, probabilities, predicted class """ if self._model is None: raise RuntimeError("Model not loaded. Call load_model() first.") if not self._warmup_done: self.warmup() model = self.model # Prepare input tensor if audio.ndim == 1: audio = audio[np.newaxis, :] # [1, T] if audio.ndim == 2: audio = audio[np.newaxis, :, :] # [1, 1, T] input_tensor = torch.from_numpy(audio).to(self._device, dtype=torch.float32) with torch.no_grad(): start = time.perf_counter() logits = model(input_tensor) inference_ms = (time.perf_counter() - start) * 1000 probs = torch.softmax(logits, dim=-1) pred_class = torch.argmax(probs, dim=-1).item() pred_prob = probs[0, pred_class].item() return { "logits": logits[0].cpu().numpy(), "probabilities": probs[0].cpu().numpy(), "predicted_class": pred_class, "predicted_label": settings.class_labels[pred_class], "confidence": pred_prob, "inference_ms": inference_ms, "is_bonafide": pred_class == 0, } def predict_batch(self, audio_batch: np.ndarray) -> list[Dict[str, Any]]: """Batch inference for multiple chunks.""" if audio_batch.ndim == 2: audio_batch = audio_batch[:, np.newaxis, :] # [B, 1, T] input_tensor = torch.from_numpy(audio_batch).to(self._device, dtype=torch.float32) model = self.model with torch.no_grad(): logits = model(input_tensor) probs = torch.softmax(logits, dim=-1) pred_classes = torch.argmax(probs, dim=-1).cpu().numpy() pred_probs = torch.max(probs, dim=-1).values.cpu().numpy() results = [] for i in range(len(audio_batch)): results.append({ "logits": logits[i].cpu().numpy(), "probabilities": probs[i].cpu().numpy(), "predicted_class": int(pred_classes[i]), "predicted_label": settings.class_labels[pred_classes[i]], "confidence": float(pred_probs[i]), "is_bonafide": pred_classes[i] == 0, }) return results def save_quantized(self, path: str) -> None: """Save quantized model state dict.""" if self._quantized_model is None: raise RuntimeError("No quantized model to save") torch.save(self._quantized_model.state_dict(), path) def get_model_info(self) -> Dict[str, Any]: """Get model metadata.""" model = self.model total_params = sum(p.numel() for p in model.parameters()) quantized_params = 0 if self.quantized: for p in model.parameters(): if p.dtype == torch.qint8: quantized_params += p.numel() return { "architecture": "RawNet1D", "quantized": self.quantized, "total_parameters": total_params, "quantized_parameters": quantized_params, "num_classes": self.config.num_classes, "class_labels": settings.class_labels, "input_shape": [1, 1, 24000], "device": str(self._device), } def create_model( model_path: str = None, quantized: bool = settings.quantized, config: RawNetConfig = None, ) -> QuantizedModelWrapper: """Factory function to create and load model.""" wrapper = QuantizedModelWrapper(config=config, quantized=quantized) wrapper.load_model(model_path) return wrapper