Download validation/app/model/quantized.py from indrajit4533/voiceguard: direct link, hf CLI and curl.
- Browser
- Download file 9.13 kB
-
https://huggingface.co/indrajit4533/voiceguard/resolve/main/validation/app/model/quantized.py
- Command line
-
hf download hf://indrajit4533/voiceguard/validation/app/model/quantized.py
-
curl -L -o quantized.py https://huggingface.co/indrajit4533/voiceguard/resolve/main/validation/app/model/quantized.py
9.13 kB
| 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 | |
| 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 |