"""Xavante - validators.py - Validadores de entrada e estado.""" from __future__ import annotations import logging from typing import Tuple import torch logger = logging.getLogger(__name__) def validate_tensor_shape(x: torch.Tensor, expected: Tuple[int, ...], name: str = "tensor") -> bool: if x.dim() != len(expected): logger.error("%s dim mismatch: got %d expected %d", name, x.dim(), len(expected)) return False for i, (got, exp) in enumerate(zip(x.shape, expected)): if exp != -1 and got != exp: logger.error("%s shape[%d] mismatch: got %d expected %d", name, i, got, exp) return False return True def validate_no_nan_inf(x: torch.Tensor, name: str = "tensor") -> bool: if torch.isnan(x).any(): logger.error("%s contains NaN", name) return False if torch.isinf(x).any(): logger.error("%s contains Inf", name) return False return True def validate_device_consistency(tensors: list, device: torch.device) -> bool: for i, t in enumerate(tensors): if t.device != device: logger.error("Tensor %d em device %s, esperado %s", i, t.device, device) return False return True def clamp_probabilities(p: torch.Tensor, eps: float = 1e-9) -> torch.Tensor: return p.clamp(min=eps, max=1.0 - eps) __all__ = [ "validate_tensor_shape", "validate_no_nan_inf", "validate_device_consistency", "clamp_probabilities", ]