DCE / streaming_intent /belief.py
That guy James Bond :)
Deploy Medical Intent Escalation API
af61b34
Raw
History Blame Contribute Delete
11.6 kB
"""
HMM Belief Updater for streaming intent classification.
Implements Markov filtering in LOG-SPACE for numerical stability:
- Predict step: log_b_pred = logsumexp(log_T + log_b_prev)
- Update step: log_b = log_b_pred + log_emission - logsumexp(...)
Supports multiple emission transforms for ablation:
- 'power': emission = (p + eps)^alpha (sharpening)
- 'temperature': emission = softmax(logits / T)
- 'isotonic': calibrated probabilities (requires calibrator)
- 'raw': pass-through neural posteriors
"""
import numpy as np
from enum import Enum
from typing import Dict, Optional, Tuple, Callable
from .config import StreamingConfig
class EmissionTransform(Enum):
"""Emission transform methods for ablation studies."""
POWER = "power" # (p + eps)^alpha sharpening
TEMPERATURE = "temperature" # softmax(logits / T)
RAW = "raw" # Pass-through neural posteriors
ISOTONIC = "isotonic" # Calibrated (requires external calibrator)
def logsumexp(log_vec: np.ndarray) -> float:
"""
Numerically stable log-sum-exp.
Edge cases:
- If all values are -inf, returns -inf (empty probability mass)
- If any value is +inf, returns +inf (handles numerical overflow gracefully)
"""
max_val = np.max(log_vec)
if np.isinf(max_val):
return max_val
return max_val + np.log(np.sum(np.exp(log_vec - max_val)))
def log_normalize(log_vec: np.ndarray) -> np.ndarray:
"""Normalize log-probabilities to sum to 1 in probability space."""
return log_vec - logsumexp(log_vec)
class BeliefUpdater:
"""
HMM-style belief state updater with log-space arithmetic.
Maintains a probability distribution over intents and updates it
using transition dynamics and emission observations from neural model.
All internal computations are done in log-space to prevent underflow
in long sequences. Probabilities are converted only for output.
Supports ablation modes:
- use_hmm=False: Raw neural posteriors (no temporal smoothing)
- emission_transform: Different emission sharpening methods
"""
def __init__(
self,
config: StreamingConfig,
use_hmm: bool = True,
emission_transform: EmissionTransform = EmissionTransform.POWER,
isotonic_calibrator: Optional[Callable] = None,
):
"""
Initialize belief updater.
Args:
config: Streaming configuration with intents, transitions, thresholds.
use_hmm: If False, bypass HMM and return raw neural posteriors.
emission_transform: Method for transforming neural outputs to emissions.
isotonic_calibrator: Optional calibrator function for ISOTONIC mode.
"""
self.config = config
self.n_intents = len(config.intents)
self.use_hmm = use_hmm
self.emission_transform = emission_transform
self.isotonic_calibrator = isotonic_calibrator
# Initialize log-space belief
prior = config.get_prior_vector()
self.log_belief = np.log(prior + config.epsilon)
self.log_belief = log_normalize(self.log_belief)
# Pre-compute log transition matrix
self.log_T = np.log(config.transition_matrix + config.epsilon)
self._step_count = 0
self._last_raw_probs: Optional[Dict[str, float]] = None
def reset(self) -> None:
"""Reset belief to prior distribution."""
prior = self.config.get_prior_vector()
self.log_belief = np.log(prior + self.config.epsilon)
self.log_belief = log_normalize(self.log_belief)
self._step_count = 0
self._last_raw_probs = None
def predict(self) -> np.ndarray:
"""
Prediction step in log-space: propagate belief through transition matrix.
log_b_pred[j] = logsumexp_i(log_T[i,j] + log_b[i])
Returns:
Log-space predicted belief state.
"""
log_b_pred = np.zeros(self.n_intents)
for j in range(self.n_intents):
# Sum over all previous states i: T[i,j] * b[i]
log_terms = self.log_T[:, j] + self.log_belief
log_b_pred[j] = logsumexp(log_terms)
return log_normalize(log_b_pred)
def compute_log_emission(
self,
neural_probs: Dict[str, float],
logits: Optional[Dict[str, float]] = None,
) -> np.ndarray:
"""
Compute log-emission likelihood from neural network outputs.
Supports multiple transform methods for ablation:
- POWER: log_emission = alpha * log(p + eps)
- TEMPERATURE: log_emission = logits / T - logsumexp(logits / T)
- RAW: log_emission = log(p + eps)
- ISOTONIC: log_emission = log(calibrate(p) + eps)
Args:
neural_probs: Dict mapping intent names to probabilities.
logits: Optional dict of raw logits for temperature scaling.
Returns:
Log-emission likelihood vector.
"""
log_emission = np.full(self.n_intents, np.log(self.config.epsilon))
if self.emission_transform == EmissionTransform.POWER:
# Power transform: sharpen probabilities
for intent, prob in neural_probs.items():
if intent in self.config.intent_to_idx:
idx = self.config.intent_to_idx[intent]
log_emission[idx] = self.config.alpha * np.log(prob + self.config.epsilon)
elif self.emission_transform == EmissionTransform.TEMPERATURE:
# Temperature scaling on logits
if logits is None:
# Fallback: invert softmax approximately
logits = {k: np.log(v + self.config.epsilon) for k, v in neural_probs.items()}
temp = self.config.temperature
scaled_logits = np.full(self.n_intents, -np.inf)
for intent, logit in logits.items():
if intent in self.config.intent_to_idx:
idx = self.config.intent_to_idx[intent]
scaled_logits[idx] = logit / temp
log_emission = log_normalize(scaled_logits)
elif self.emission_transform == EmissionTransform.RAW:
# Pass-through: no sharpening
for intent, prob in neural_probs.items():
if intent in self.config.intent_to_idx:
idx = self.config.intent_to_idx[intent]
log_emission[idx] = np.log(prob + self.config.epsilon)
elif self.emission_transform == EmissionTransform.ISOTONIC:
# Isotonic calibration
if self.isotonic_calibrator is not None:
calibrated = self.isotonic_calibrator(neural_probs)
for intent, prob in calibrated.items():
if intent in self.config.intent_to_idx:
idx = self.config.intent_to_idx[intent]
log_emission[idx] = np.log(prob + self.config.epsilon)
else:
# Fallback to raw if no calibrator
for intent, prob in neural_probs.items():
if intent in self.config.intent_to_idx:
idx = self.config.intent_to_idx[intent]
log_emission[idx] = np.log(prob + self.config.epsilon)
return log_emission
def update(
self,
neural_probs: Dict[str, float],
logits: Optional[Dict[str, float]] = None,
) -> np.ndarray:
"""
Full belief update in log-space.
If use_hmm=True:
1. log_b_pred = predict() (transition dynamics)
2. log_emission = compute_log_emission(neural_probs)
3. log_b = normalize(log_b_pred + log_emission)
If use_hmm=False (ablation mode):
Directly use neural posteriors as belief.
Args:
neural_probs: Dict mapping intent names to probabilities from neural model.
logits: Optional dict of raw logits for temperature scaling.
Returns:
Updated belief state (in probability space).
"""
self._last_raw_probs = neural_probs.copy()
if not self.use_hmm:
# Ablation: bypass HMM, use raw neural posteriors
for intent, prob in neural_probs.items():
if intent in self.config.intent_to_idx:
idx = self.config.intent_to_idx[intent]
self.log_belief[idx] = np.log(prob + self.config.epsilon)
self.log_belief = log_normalize(self.log_belief)
self._step_count += 1
return self.belief
# Predict step
log_b_pred = self.predict()
# Compute log-emission likelihood
log_emission = self.compute_log_emission(neural_probs, logits)
# Update step: element-wise addition in log-space and normalize
log_b_updated = log_b_pred + log_emission
self.log_belief = log_normalize(log_b_updated)
self._step_count += 1
return self.belief
def update_with_log_emission(self, log_emission: np.ndarray) -> np.ndarray:
"""
Update belief with pre-computed log-emission vector.
Args:
log_emission: Pre-computed log-emission likelihood vector.
Returns:
Updated belief state (in probability space).
"""
if not self.use_hmm:
self.log_belief = log_normalize(log_emission)
self._step_count += 1
return self.belief
log_b_pred = self.predict()
log_b_updated = log_b_pred + log_emission
self.log_belief = log_normalize(log_b_updated)
self._step_count += 1
return self.belief
@property
def belief(self) -> np.ndarray:
"""Get current belief in probability space."""
return np.exp(self.log_belief)
def get_belief_dict(self) -> Dict[str, float]:
"""Get current belief as dictionary."""
probs = self.belief
return {
self.config.idx_to_intent[i]: float(probs[i])
for i in range(self.n_intents)
}
def get_top_intent(self) -> Tuple[str, float]:
"""Get intent with highest belief probability."""
probs = self.belief
idx = int(np.argmax(probs))
return self.config.idx_to_intent[idx], float(probs[idx])
def get_intent_prob(self, intent: str) -> float:
"""Get belief probability for specific intent."""
if intent in self.config.intent_to_idx:
idx = self.config.intent_to_idx[intent]
return float(np.exp(self.log_belief[idx]))
return 0.0
def get_log_belief(self) -> np.ndarray:
"""Get current log-belief (for debugging/analysis)."""
return self.log_belief.copy()
def get_last_raw_probs(self) -> Optional[Dict[str, float]]:
"""Get last raw neural probabilities (before HMM filtering)."""
return self._last_raw_probs
@property
def step_count(self) -> int:
"""Number of update steps performed."""
return self._step_count
def is_normalized(self, tol: float = 1e-6) -> bool:
"""Check if belief state is properly normalized."""
return abs(self.belief.sum() - 1.0) < tol