File size: 2,177 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
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
"""Xavante - tensor_ops.py - Operacoes tensoriais otimizadas."""
from __future__ import annotations

import logging
import math
from typing import Tuple

import torch

logger = logging.getLogger(__name__)


def safe_softmax(x: torch.Tensor, dim: int = -1) -> torch.Tensor:
    """Softmax com subtracao do max para estabilidade numerica."""
    return torch.softmax(x - x.amax(dim=dim, keepdim=True), dim=dim)


def safe_log(x: torch.Tensor, eps: float = 1e-9) -> torch.Tensor:
    return torch.log(x.clamp_min(eps))


def cosine_similarity(a: torch.Tensor, b: torch.Tensor, dim: int = -1, eps: float = 1e-8) -> torch.Tensor:
    return torch.nn.functional.cosine_similarity(a, b, dim=dim, eps=eps)


def orthogonal_init(weight: torch.Tensor) -> None:
    """Inicializa weight como matriz ortogonal."""
    nn_init = torch.nn.init
    nn_init.orthogonal_(weight)


def skew_to_rotation(skew: torch.Tensor) -> torch.Tensor:
    """Converte matriz skew-symmetric para rotation via Rodrigues."""
    n = skew.shape[0]
    I = torch.eye(n, device=skew.device, dtype=skew.dtype)
    A = skew
    A2 = A @ A
    # Exponential via series truncada
    return I + A + 0.5 * A2 + (1.0 / 6.0) * (A2 @ A)


def chunked_matmul(a: torch.Tensor, b: torch.Tensor, chunk: int = 4096) -> torch.Tensor:
    """Matmul em chunks para evitar OOM com matrizes grandes."""
    if a.shape[1] <= chunk:
        return a @ b
    out_chunks = []
    for i in range(0, a.shape[0], chunk):
        out_chunks.append(a[i : i + chunk] @ b)
    return torch.cat(out_chunks, dim=0)


def padding_mask(lengths: torch.Tensor, max_len: int) -> torch.Tensor:
    """Cria mascara bool [B, max_len] onde True = token valido."""
    idx = torch.arange(max_len, device=lengths.device).unsqueeze(0)
    return idx < lengths.unsqueeze(1)


def causal_mask(seq_len: int, device: torch.device) -> torch.Tensor:
    """Mascara causal [seq_len, seq_len] (1 = permitido)."""
    return torch.tril(torch.ones(seq_len, seq_len, device=device))


__all__ = [
    "safe_softmax",
    "safe_log",
    "cosine_similarity",
    "orthogonal_init",
    "skew_to_rotation",
    "chunked_matmul",
    "padding_mask",
    "causal_mask",
]