Download src/bigru_t/utils/validators.py from PowerMachine/BiGRU_T_version: direct link, hf CLI and curl.
- Browser
- Download file 1.48 kB
-
https://huggingface.co/PowerMachine/BiGRU_T_version/resolve/main/src/bigru_t/utils/validators.py
- Command line
-
hf download hf://PowerMachine/BiGRU_T_version/src/bigru_t/utils/validators.py
-
curl -L -o validators.py https://huggingface.co/PowerMachine/BiGRU_T_version/resolve/main/src/bigru_t/utils/validators.py
1.48 kB
| """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", | |
| ] | |