voiceguard / validation /app /model /quantized.py
indrajit4533's picture
Add reproducible validation evidence kit (harness + results + README)
c863b63 verified
Raw History Blame Contribute Delete
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
@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