BiGRU_T_version / src /bigru_t /utils /tensor_ops.py
PowerMachine's picture
Upload folder using huggingface_hub
3275441 verified
Raw History Blame Contribute Delete
2.18 kB
"""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",
]