| """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", |
| ] |
|
|