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