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