File size: 1,478 Bytes
3275441 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 | """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",
]
|